nullcal 0.2.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- nullcal/__init__.py +16 -0
- nullcal/calibration.py +189 -0
- nullcal/clustering/__init__.py +0 -0
- nullcal/clustering/base.py +39 -0
- nullcal/clustering/injection.py +63 -0
- nullcal/clustering/precompute.py +30 -0
- nullcal/clustering/single.py +208 -0
- nullcal/clustering/time_frequency_map.py +42 -0
- nullcal/data.py +93 -0
- nullcal/likelihood/__init__.py +9 -0
- nullcal/likelihood/recalibration_likelihood.py +227 -0
- nullcal/metadata/__init__.py +0 -0
- nullcal/metadata/yaml.py +30 -0
- nullcal/null_stream/__init__.py +0 -0
- nullcal/null_stream/calibration.py +40 -0
- nullcal/null_stream/null_stream.py +252 -0
- nullcal/null_stream/projector.py +48 -0
- nullcal/null_stream/whiten.py +106 -0
- nullcal/result/__init__.py +5 -0
- nullcal/result/result.py +50 -0
- nullcal/result/utils.py +13 -0
- nullcal/sampler.py +113 -0
- nullcal/studies/__init__.py +1 -0
- nullcal/studies/lwa_leakage.py +513 -0
- nullcal/studies/spline_resolution.py +148 -0
- nullcal/time_frequency_transform/README.md +7 -0
- nullcal/time_frequency_transform/__init__.py +23 -0
- nullcal/time_frequency_transform/inverse_wavelet_freq_funcs.py +42 -0
- nullcal/time_frequency_transform/inverse_wavelet_time_funcs.py +49 -0
- nullcal/time_frequency_transform/stft.py +45 -0
- nullcal/time_frequency_transform/transform_freq_funcs.py +180 -0
- nullcal/time_frequency_transform/transform_time_funcs.py +60 -0
- nullcal/time_frequency_transform/utils.py +21 -0
- nullcal/time_frequency_transform/wavelet_transforms.py +249 -0
- nullcal/utils/__init__.py +0 -0
- nullcal/utils/log.py +72 -0
- nullcal/utils/snr.py +29 -0
- nullcal/version.py +9 -0
- nullcal-0.2.0.dist-info/METADATA +222 -0
- nullcal-0.2.0.dist-info/RECORD +42 -0
- nullcal-0.2.0.dist-info/WHEEL +4 -0
- nullcal-0.2.0.dist-info/licenses/LICENSE +21 -0
|
@@ -0,0 +1,227 @@
|
|
|
1
|
+
"""Pure JAX log density for null-stream recalibration inference."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from collections.abc import Mapping
|
|
6
|
+
|
|
7
|
+
import jax
|
|
8
|
+
|
|
9
|
+
jax.config.update("jax_enable_x64", True)
|
|
10
|
+
|
|
11
|
+
import jax.numpy as jnp # noqa: E402
|
|
12
|
+
import numpy as np # noqa: E402
|
|
13
|
+
|
|
14
|
+
from ..calibration import MINIMUM_CUBIC_SPLINE_KNOTS, calibration_factor, calibration_log_prior # noqa: E402
|
|
15
|
+
from ..data import InterferometerData # noqa: E402
|
|
16
|
+
from ..null_stream.whiten import ( # noqa: E402
|
|
17
|
+
compute_whitened_antenna_response,
|
|
18
|
+
compute_whitened_frequency_domain_strain,
|
|
19
|
+
)
|
|
20
|
+
from ..time_frequency_transform.transform_freq_funcs import ( # noqa: E402
|
|
21
|
+
phitilde_vec_norm,
|
|
22
|
+
transform_wavelet_freq_helper,
|
|
23
|
+
)
|
|
24
|
+
from ..time_frequency_transform.wavelet_transforms import WaveletTransform # noqa: E402
|
|
25
|
+
|
|
26
|
+
PARAMETER_NAMES = frozenset({"amplitude", "phase"})
|
|
27
|
+
PARAMETER_ARRAY_NDIM = 2
|
|
28
|
+
ET_BEAM_PATTERN = np.array(
|
|
29
|
+
[[-1.0 / np.sqrt(6.0), -1.0 / np.sqrt(2.0)], [np.sqrt(6.0) / 3.0, 0.0], [-1.0 / np.sqrt(6.0), 1.0 / np.sqrt(2.0)]]
|
|
30
|
+
)
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
def _broadcast_prior(value, shape: tuple[int, int], name: str) -> jax.Array:
|
|
34
|
+
array = jnp.asarray(value, dtype=jnp.float64)
|
|
35
|
+
try:
|
|
36
|
+
return jnp.broadcast_to(array, shape)
|
|
37
|
+
except ValueError as error:
|
|
38
|
+
raise ValueError(f"{name} must be scalar or broadcast to detector-by-knot shape {shape}") from error
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
def _validated_knots(knot_frequencies, detector_count: int) -> np.ndarray:
|
|
42
|
+
knots = np.asarray(knot_frequencies, dtype=np.float64)
|
|
43
|
+
if knots.ndim == 1:
|
|
44
|
+
knots = np.broadcast_to(knots, (detector_count, knots.size)).copy()
|
|
45
|
+
if knots.ndim != PARAMETER_ARRAY_NDIM or knots.shape[0] != detector_count:
|
|
46
|
+
raise ValueError("knot_frequencies must have shape (knot,) or (detector, knot)")
|
|
47
|
+
if knots.shape[1] < MINIMUM_CUBIC_SPLINE_KNOTS:
|
|
48
|
+
raise ValueError("a cubic spline requires at least four knots")
|
|
49
|
+
if np.any(knots <= 0.0) or np.any(np.diff(knots, axis=1) <= 0.0):
|
|
50
|
+
raise ValueError("knot_frequencies must be positive and strictly increasing")
|
|
51
|
+
return knots
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
def _whitened_inputs(interferometers, shared_mask, supplied_whitened_strain):
|
|
55
|
+
delta_f = 1.0 / float(interferometers.duration)
|
|
56
|
+
psd = np.asarray(interferometers.psd)
|
|
57
|
+
whitened_response = compute_whitened_antenna_response(ET_BEAM_PATTERN, psd, delta_f, shared_mask)
|
|
58
|
+
if supplied_whitened_strain is None:
|
|
59
|
+
whitened_strain = compute_whitened_frequency_domain_strain(
|
|
60
|
+
np.asarray(interferometers.strain), psd, delta_f, shared_mask
|
|
61
|
+
)
|
|
62
|
+
else:
|
|
63
|
+
whitened_strain = np.asarray(supplied_whitened_strain)
|
|
64
|
+
if whitened_strain.shape != psd.shape:
|
|
65
|
+
raise ValueError("whitened_frequency_domain_strain must match the detector data shape")
|
|
66
|
+
if not np.all(whitened_strain[:, ~shared_mask] == 0.0):
|
|
67
|
+
raise ValueError("whitened_frequency_domain_strain must be zero outside the shared mask")
|
|
68
|
+
return psd, whitened_response, whitened_strain
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
class RecalibrationLikelihood:
|
|
72
|
+
"""Immutable-data recalibration posterior with a pure ``logdensity_fn``.
|
|
73
|
+
|
|
74
|
+
One instance represents one detector-network realization. Parameters are a
|
|
75
|
+
pytree with ``amplitude`` and ``phase`` arrays of shape ``(detector, knot)``.
|
|
76
|
+
Knot frequencies and Gaussian prior hyperparameters are fixed model data,
|
|
77
|
+
not sampled coordinates.
|
|
78
|
+
"""
|
|
79
|
+
|
|
80
|
+
def __init__(
|
|
81
|
+
self,
|
|
82
|
+
interferometers: InterferometerData,
|
|
83
|
+
knot_frequencies,
|
|
84
|
+
*,
|
|
85
|
+
time_frequency_filter: np.ndarray,
|
|
86
|
+
wavelet_transform_frequency_resolution: float = 4.0,
|
|
87
|
+
wavelet_transform_nx: float = 4.0,
|
|
88
|
+
amplitude_prior_mean=0.0,
|
|
89
|
+
amplitude_prior_sigma=0.05,
|
|
90
|
+
phase_prior_mean=0.0,
|
|
91
|
+
phase_prior_sigma=0.05,
|
|
92
|
+
whitened_frequency_domain_strain=None,
|
|
93
|
+
) -> None:
|
|
94
|
+
if not isinstance(interferometers, InterferometerData):
|
|
95
|
+
raise TypeError("interferometers must be an InterferometerData instance")
|
|
96
|
+
if len(interferometers) != ET_BEAM_PATTERN.shape[0]:
|
|
97
|
+
raise ValueError("the recalibration likelihood currently requires the three-detector ET triangle")
|
|
98
|
+
|
|
99
|
+
transform = WaveletTransform(
|
|
100
|
+
duration=float(interferometers.duration),
|
|
101
|
+
sampling_frequency=float(interferometers.sampling_frequency),
|
|
102
|
+
frequency_resolution=wavelet_transform_frequency_resolution,
|
|
103
|
+
nx=wavelet_transform_nx,
|
|
104
|
+
)
|
|
105
|
+
time_frequency_filter = np.asarray(time_frequency_filter, dtype=bool)
|
|
106
|
+
if time_frequency_filter.shape != transform.shape:
|
|
107
|
+
raise ValueError(
|
|
108
|
+
f"time_frequency_filter has shape {time_frequency_filter.shape}; expected {transform.shape}"
|
|
109
|
+
)
|
|
110
|
+
|
|
111
|
+
shared_mask = np.all(np.asarray(interferometers.mask, dtype=bool), axis=0)
|
|
112
|
+
frequency_indices = np.flatnonzero(shared_mask)
|
|
113
|
+
if frequency_indices.size == 0:
|
|
114
|
+
raise ValueError("the shared detector frequency mask must not be empty")
|
|
115
|
+
masked_frequencies = np.asarray(interferometers.frequency_array)[frequency_indices]
|
|
116
|
+
if np.any(masked_frequencies <= 0.0):
|
|
117
|
+
raise ValueError("calibration frequencies must be positive")
|
|
118
|
+
|
|
119
|
+
knots = _validated_knots(knot_frequencies, len(interferometers))
|
|
120
|
+
psd, whitened_response, whitened_strain = _whitened_inputs(
|
|
121
|
+
interferometers, shared_mask, whitened_frequency_domain_strain
|
|
122
|
+
)
|
|
123
|
+
|
|
124
|
+
parameter_shape = knots.shape
|
|
125
|
+
self.interferometers = interferometers
|
|
126
|
+
self.knot_frequencies = jnp.asarray(knots)
|
|
127
|
+
self.time_frequency_filter = jnp.asarray(time_frequency_filter)
|
|
128
|
+
self.time_frequency_transform = transform
|
|
129
|
+
self.amplitude_prior_mean = _broadcast_prior(amplitude_prior_mean, parameter_shape, "amplitude_prior_mean")
|
|
130
|
+
self.amplitude_prior_sigma = _broadcast_prior(amplitude_prior_sigma, parameter_shape, "amplitude_prior_sigma")
|
|
131
|
+
self.phase_prior_mean = _broadcast_prior(phase_prior_mean, parameter_shape, "phase_prior_mean")
|
|
132
|
+
self.phase_prior_sigma = _broadcast_prior(phase_prior_sigma, parameter_shape, "phase_prior_sigma")
|
|
133
|
+
if np.any(np.asarray(self.amplitude_prior_sigma) <= 0.0) or np.any(np.asarray(self.phase_prior_sigma) <= 0.0):
|
|
134
|
+
raise ValueError("prior standard deviations must be positive")
|
|
135
|
+
|
|
136
|
+
self._frequency_indices = jnp.asarray(frequency_indices)
|
|
137
|
+
self._masked_frequencies = jnp.asarray(masked_frequencies)
|
|
138
|
+
self._whitened_antenna_response = jnp.asarray(whitened_response[frequency_indices])
|
|
139
|
+
self._whitened_frequency_domain_strain = jnp.asarray(whitened_strain[:, frequency_indices])
|
|
140
|
+
self._frequency_count = psd.shape[1]
|
|
141
|
+
n_time, n_frequency = transform.shape
|
|
142
|
+
self._wavelet_n_time = n_time
|
|
143
|
+
self._wavelet_n_frequency = n_frequency
|
|
144
|
+
self._wavelet_window = (2.0 / n_frequency) * phitilde_vec_norm(n_frequency, n_time, wavelet_transform_nx)
|
|
145
|
+
self._wavelet_scale = jnp.sqrt((self._frequency_count - 1) * 2.0)
|
|
146
|
+
|
|
147
|
+
@property
|
|
148
|
+
def parameter_shape(self) -> tuple[int, int]:
|
|
149
|
+
"""Required shape of each parameter array."""
|
|
150
|
+
return self.knot_frequencies.shape
|
|
151
|
+
|
|
152
|
+
def _validated_parameters(self, params: Mapping[str, jax.Array]) -> tuple[jax.Array, jax.Array]:
|
|
153
|
+
if not isinstance(params, Mapping) or set(params) != PARAMETER_NAMES:
|
|
154
|
+
raise ValueError("params must contain exactly amplitude and phase")
|
|
155
|
+
amplitude = jnp.asarray(params["amplitude"], dtype=jnp.float64)
|
|
156
|
+
phase = jnp.asarray(params["phase"], dtype=jnp.float64)
|
|
157
|
+
if amplitude.shape != self.parameter_shape or phase.shape != self.parameter_shape:
|
|
158
|
+
raise ValueError(
|
|
159
|
+
f"amplitude and phase must each have shape {self.parameter_shape}; "
|
|
160
|
+
f"received {amplitude.shape} and {phase.shape}"
|
|
161
|
+
)
|
|
162
|
+
return amplitude, phase
|
|
163
|
+
|
|
164
|
+
def _calibration_factor(self, amplitude: jax.Array, phase: jax.Array) -> jax.Array:
|
|
165
|
+
return jax.vmap(calibration_factor, in_axes=(None, 0, 0, 0))(
|
|
166
|
+
self._masked_frequencies,
|
|
167
|
+
self.knot_frequencies,
|
|
168
|
+
amplitude,
|
|
169
|
+
phase,
|
|
170
|
+
)
|
|
171
|
+
|
|
172
|
+
def _frequency_domain_null_stream(self, amplitude: jax.Array, phase: jax.Array) -> jax.Array:
|
|
173
|
+
factors = self._calibration_factor(amplitude, phase)
|
|
174
|
+
response = self._whitened_antenna_response * jnp.swapaxes(factors, 0, 1)[:, :, None]
|
|
175
|
+
response_dagger = jnp.swapaxes(jnp.conj(response), 1, 2)
|
|
176
|
+
gram = response_dagger @ response
|
|
177
|
+
projected_response = response @ jnp.linalg.solve(gram, response_dagger)
|
|
178
|
+
projector = jnp.eye(response.shape[1], dtype=response.dtype)[None, :, :] - projected_response
|
|
179
|
+
masked_null_stream = jnp.einsum("fij,jf->if", projector, self._whitened_frequency_domain_strain)
|
|
180
|
+
return (
|
|
181
|
+
jnp.zeros((self.parameter_shape[0], self._frequency_count), dtype=masked_null_stream.dtype)
|
|
182
|
+
.at[:, self._frequency_indices]
|
|
183
|
+
.set(masked_null_stream)
|
|
184
|
+
)
|
|
185
|
+
|
|
186
|
+
def _time_frequency_null_stream(self, amplitude: jax.Array, phase: jax.Array) -> jax.Array:
|
|
187
|
+
frequency_domain = self._frequency_domain_null_stream(amplitude, phase)
|
|
188
|
+
|
|
189
|
+
def transform(detector_data):
|
|
190
|
+
return (
|
|
191
|
+
transform_wavelet_freq_helper(
|
|
192
|
+
detector_data,
|
|
193
|
+
self._wavelet_n_frequency,
|
|
194
|
+
self._wavelet_n_time,
|
|
195
|
+
self._wavelet_window,
|
|
196
|
+
)
|
|
197
|
+
* self._wavelet_scale
|
|
198
|
+
)
|
|
199
|
+
|
|
200
|
+
time_frequency = jax.vmap(transform)(frequency_domain)
|
|
201
|
+
return jnp.where(self.time_frequency_filter[None, :, :], time_frequency, 0.0)
|
|
202
|
+
|
|
203
|
+
def log_likelihood_fn(self, params: Mapping[str, jax.Array]) -> jax.Array:
|
|
204
|
+
"""Return the pure JAX null-stream log likelihood for ``params``."""
|
|
205
|
+
amplitude, phase = self._validated_parameters(params)
|
|
206
|
+
null_stream = self._time_frequency_null_stream(amplitude, phase)
|
|
207
|
+
return -0.5 * jnp.sum(jnp.abs(null_stream) ** 2)
|
|
208
|
+
|
|
209
|
+
def logdensity_fn(self, params: Mapping[str, jax.Array]) -> jax.Array:
|
|
210
|
+
"""Return normalized Gaussian log prior plus null-stream log likelihood."""
|
|
211
|
+
amplitude, phase = self._validated_parameters(params)
|
|
212
|
+
return self.log_likelihood_fn(params) + calibration_log_prior(
|
|
213
|
+
amplitude,
|
|
214
|
+
phase,
|
|
215
|
+
self.amplitude_prior_mean,
|
|
216
|
+
self.amplitude_prior_sigma,
|
|
217
|
+
self.phase_prior_mean,
|
|
218
|
+
self.phase_prior_sigma,
|
|
219
|
+
)
|
|
220
|
+
|
|
221
|
+
def noise_log_likelihood(self) -> jax.Array:
|
|
222
|
+
"""Return the likelihood for an exactly uncalibrated response."""
|
|
223
|
+
zeros = jnp.zeros(self.parameter_shape, dtype=jnp.float64)
|
|
224
|
+
return self.log_likelihood_fn({"amplitude": zeros, "phase": zeros})
|
|
225
|
+
|
|
226
|
+
|
|
227
|
+
__all__ = ["RecalibrationLikelihood"]
|
|
File without changes
|
nullcal/metadata/yaml.py
ADDED
|
@@ -0,0 +1,30 @@
|
|
|
1
|
+
"""YAML read and write utilities."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import yaml
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
def write_to_yaml(fname: str, data: dict) -> None:
|
|
9
|
+
"""Write a dictionary to a YAML file.
|
|
10
|
+
|
|
11
|
+
Args:
|
|
12
|
+
fname (str): File name.
|
|
13
|
+
data (dict): A dictionary of data.
|
|
14
|
+
"""
|
|
15
|
+
with open(fname, mode="w", encoding="utf-8") as f:
|
|
16
|
+
yaml.safe_dump(data, f, default_flow_style=False)
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def load_from_yaml(fname: str) -> dict:
|
|
20
|
+
"""Load from a YAML file.
|
|
21
|
+
|
|
22
|
+
Args:
|
|
23
|
+
fname (str): File name.
|
|
24
|
+
|
|
25
|
+
Returns:
|
|
26
|
+
dict: A dictionary of data.
|
|
27
|
+
"""
|
|
28
|
+
with open(fname, encoding="utf-8") as f:
|
|
29
|
+
data = yaml.safe_load(f)
|
|
30
|
+
return data
|
|
File without changes
|
|
@@ -0,0 +1,40 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Functions for including calibration factors into
|
|
3
|
+
the antenna response function.
|
|
4
|
+
"""
|
|
5
|
+
|
|
6
|
+
from __future__ import annotations
|
|
7
|
+
|
|
8
|
+
import jax
|
|
9
|
+
|
|
10
|
+
# Calibration inference is numerically unstable in JAX's default float32 mode.
|
|
11
|
+
# Keep this kernel consistent with nullcal.calibration, which establishes x64 as
|
|
12
|
+
# a process-wide requirement for the calibration path.
|
|
13
|
+
jax.config.update("jax_enable_x64", True)
|
|
14
|
+
|
|
15
|
+
import jax.numpy as jnp # noqa: E402
|
|
16
|
+
import numpy as np # noqa: E402
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def compute_calibrated_whitened_antenna_response(
|
|
20
|
+
whitened_antenna_response: np.ndarray, calibration_factor: np.ndarray, frequency_mask: np.ndarray
|
|
21
|
+
) -> np.ndarray:
|
|
22
|
+
"""Compute the whitened antenna response function
|
|
23
|
+
with the calibration factor included.
|
|
24
|
+
|
|
25
|
+
Args:
|
|
26
|
+
whitened_antenna_response (np.ndarray): Whitened antenna response function.
|
|
27
|
+
Dimensions: (frequency, detector, polarization).
|
|
28
|
+
calibration_factor (np.ndarray): Calibration factor.
|
|
29
|
+
Dimensions: (detector, frequency).
|
|
30
|
+
frequency_mask (np.ndarray): Frequency mask.
|
|
31
|
+
Dimensions: (frequency,)
|
|
32
|
+
|
|
33
|
+
Returns:
|
|
34
|
+
np.ndarray: Calibrated antenna response function.
|
|
35
|
+
"""
|
|
36
|
+
calibration_factor = jnp.asarray(calibration_factor)
|
|
37
|
+
response = jnp.asarray(whitened_antenna_response, dtype=calibration_factor.dtype)
|
|
38
|
+
mask = jnp.asarray(frequency_mask, dtype=bool)[:, None, None]
|
|
39
|
+
calibrated = response * jnp.swapaxes(calibration_factor, 0, 1)[:, :, None]
|
|
40
|
+
return jnp.where(mask, calibrated, jnp.zeros((), dtype=calibrated.dtype))
|
|
@@ -0,0 +1,252 @@
|
|
|
1
|
+
"""A submodule for null stream calculation."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
|
|
7
|
+
from ..calibration import MINIMUM_CUBIC_SPLINE_KNOTS, calibration_factor
|
|
8
|
+
from ..data import InterferometerData
|
|
9
|
+
from ..time_frequency_transform.wavelet_transforms import WaveletTransform
|
|
10
|
+
from .calibration import compute_calibrated_whitened_antenna_response
|
|
11
|
+
from .projector import compute_projector
|
|
12
|
+
from .whiten import compute_whitened_antenna_response, compute_whitened_frequency_domain_strain
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def compute_projected_strain_data(
|
|
16
|
+
projector: np.ndarray, strain_data: np.ndarray, frequency_mask: np.ndarray
|
|
17
|
+
) -> np.ndarray:
|
|
18
|
+
"""Compute the projected strain data.
|
|
19
|
+
|
|
20
|
+
Args:
|
|
21
|
+
projector (np.ndarray): Projector. Dimensions: (frequency, detector, detector).
|
|
22
|
+
strain_data (np.ndarray): Strain data. Dimensions: (detector, frequency).
|
|
23
|
+
frequency_mask (np.ndarray): Frequency mask. Dimensions: (frequency,).
|
|
24
|
+
|
|
25
|
+
Raises:
|
|
26
|
+
ValueError: Projector shape mismatch. project must have the same dimensions in the two last axes.
|
|
27
|
+
ValueError: Shape mismatch. projector and strain_data must have the same detector and frequency dimensions.
|
|
28
|
+
ValueError: Shape mismatch. strain_data and frequency_mask must have the same frequency dimension.
|
|
29
|
+
|
|
30
|
+
Returns:
|
|
31
|
+
np.ndarray: Projected strain data. Dimensions: (detector, frequency).
|
|
32
|
+
"""
|
|
33
|
+
n_freq_1, n_det_1, n_det_2 = projector.shape
|
|
34
|
+
if n_det_1 != n_det_2:
|
|
35
|
+
raise ValueError(
|
|
36
|
+
"Shape mismatch."
|
|
37
|
+
f"projector: (frequency={n_freq_1},detector={n_det_1}, detector={n_det_2})."
|
|
38
|
+
"project must have the same dimensions in the two last axes."
|
|
39
|
+
)
|
|
40
|
+
n_det_3, n_freq_2 = strain_data.shape
|
|
41
|
+
if n_det_1 != n_det_3 or n_freq_1 != n_freq_2:
|
|
42
|
+
raise ValueError(
|
|
43
|
+
"Shape mismatch."
|
|
44
|
+
f"projector: (frequency={n_freq_1},detector={n_det_1}, detector={n_det_2})."
|
|
45
|
+
f"strain_data: (detector={n_det_3},frequency={n_freq_2})."
|
|
46
|
+
"projector and strain_data must have the same detector and frequency dimensions."
|
|
47
|
+
)
|
|
48
|
+
n_freq_3 = frequency_mask.shape[0]
|
|
49
|
+
if n_freq_2 != n_freq_3:
|
|
50
|
+
raise ValueError(
|
|
51
|
+
"Shape mismatch."
|
|
52
|
+
f"strain_data: (detector={n_det_3},frequency={n_freq_2})."
|
|
53
|
+
f"frequency_mask: (frequency={n_freq_3})."
|
|
54
|
+
"strain_data and frequency_mask must have the same frequency dimension."
|
|
55
|
+
)
|
|
56
|
+
output = np.zeros_like(strain_data)
|
|
57
|
+
output[:, frequency_mask] = np.einsum("fij,jf->if", projector[frequency_mask, :, :], strain_data[:, frequency_mask])
|
|
58
|
+
return output
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
class NullStream:
|
|
62
|
+
"""A class to handle null stream calculation."""
|
|
63
|
+
|
|
64
|
+
def __init__(
|
|
65
|
+
self,
|
|
66
|
+
interferometers: InterferometerData,
|
|
67
|
+
time_frequency_transform: WaveletTransform,
|
|
68
|
+
time_frequency_filter: np.ndarray,
|
|
69
|
+
):
|
|
70
|
+
"""A null stream calculator.
|
|
71
|
+
|
|
72
|
+
Args:
|
|
73
|
+
interferometers (InterferometerData): Frozen detector arrays and metadata.
|
|
74
|
+
time_frequency_transform (WaveletTransform): A WaveletTransform instance.
|
|
75
|
+
time_frequency_filter (np.ndarray): A time-frequency filter.
|
|
76
|
+
"""
|
|
77
|
+
self.interferometers = interferometers
|
|
78
|
+
self.time_frequency_transform = time_frequency_transform
|
|
79
|
+
self.time_frequency_filter = time_frequency_filter
|
|
80
|
+
|
|
81
|
+
# Pre-compute the whitened quantities.
|
|
82
|
+
self.frequency_mask = np.all(self.interferometers.mask, axis=0)
|
|
83
|
+
self.masked_frequency_array = self.interferometers.frequency_array[self.frequency_mask]
|
|
84
|
+
# Construct the noise weighed antenna pattern
|
|
85
|
+
# This is the orthogonalized beam pattern matrix, correct for ET only,
|
|
86
|
+
# ignoring the small difference in location of the detectors.
|
|
87
|
+
beam_pattern_matrix = np.array(
|
|
88
|
+
[[-1.0 / np.sqrt(6), -1 / np.sqrt(2)], [np.sqrt(6) / 3, 0], [-1 / np.sqrt(6), 1 / np.sqrt(2)]]
|
|
89
|
+
)
|
|
90
|
+
power_spectral_density_array = np.asarray(interferometers.psd)
|
|
91
|
+
self._whitened_antenna_response = compute_whitened_antenna_response(
|
|
92
|
+
beam_pattern_matrix,
|
|
93
|
+
power_spectral_density_array,
|
|
94
|
+
1 / self.interferometers.duration,
|
|
95
|
+
self.frequency_mask,
|
|
96
|
+
)
|
|
97
|
+
self._whitened_frequency_domain_strain_array = compute_whitened_frequency_domain_strain(
|
|
98
|
+
frequency_domain_strain_array=np.asarray(interferometers.strain),
|
|
99
|
+
power_spectral_density_array=power_spectral_density_array,
|
|
100
|
+
delta_f=1.0 / interferometers.duration,
|
|
101
|
+
frequency_mask=self.frequency_mask,
|
|
102
|
+
)
|
|
103
|
+
|
|
104
|
+
def compute_uncalibrated_frequency_domain_null_stream(self) -> np.ndarray:
|
|
105
|
+
"""Compute the uncalibrated frequency domain null stream.
|
|
106
|
+
|
|
107
|
+
Returns:
|
|
108
|
+
np.ndarray: Uncalibrated frequency domain null stream. Dimensions: (detector, frequency).
|
|
109
|
+
"""
|
|
110
|
+
# Dimensions: (frequency, detector, detector)
|
|
111
|
+
projector = compute_projector(self._whitened_antenna_response, frequency_mask=self.frequency_mask)
|
|
112
|
+
# Dimensions: (frequency, detector)
|
|
113
|
+
|
|
114
|
+
return np.einsum("ijk,ki->ji", projector, self._whitened_frequency_domain_strain_array)
|
|
115
|
+
|
|
116
|
+
def compute_calibrated_frequency_domain_null_stream(self, calibration_factor: np.ndarray) -> np.ndarray:
|
|
117
|
+
"""Compute the calibrated frequency domain null stream.
|
|
118
|
+
|
|
119
|
+
Args:
|
|
120
|
+
calibration_factor (np.ndarray): Calibration factor. Dimensions: (detector, frequency).
|
|
121
|
+
|
|
122
|
+
Returns:
|
|
123
|
+
np.ndarray: Calibrated frequency domain null stream. Dimensions: (detector, frequency).
|
|
124
|
+
"""
|
|
125
|
+
calibrated_whitened_antenna_response = compute_calibrated_whitened_antenna_response(
|
|
126
|
+
self._whitened_antenna_response, calibration_factor, self.frequency_mask
|
|
127
|
+
)
|
|
128
|
+
projector = compute_projector(calibrated_whitened_antenna_response, frequency_mask=self.frequency_mask)
|
|
129
|
+
# Dimensions: (frequency, detector)
|
|
130
|
+
|
|
131
|
+
return np.einsum("ijk,ki->ji", projector, self._whitened_frequency_domain_strain_array)
|
|
132
|
+
|
|
133
|
+
def compute_uncalibrated_time_frequency_domain_null_stream(self) -> np.ndarray:
|
|
134
|
+
"""Compute the uncalibrated time-frequency domain null stream, confined to the filter.
|
|
135
|
+
|
|
136
|
+
The returned array is zero outside ``time_frequency_filter``, matching what
|
|
137
|
+
``compute_calibrated_time_frequency_domain_null_stream_from_parameters`` returns. That
|
|
138
|
+
symmetry is the point: ``noise_log_likelihood`` sums this array's energy and
|
|
139
|
+
``log_likelihood`` sums the calibrated one's, so if the two were confined to different
|
|
140
|
+
domains their difference — the log Bayes factor, and anything derived from it — would be
|
|
141
|
+
normalised against different sets of pixels.
|
|
142
|
+
|
|
143
|
+
Returns:
|
|
144
|
+
np.ndarray: Uncalibrated time-frequency domain null stream, zero outside the filter.
|
|
145
|
+
"""
|
|
146
|
+
uncalibrated_frequency_domain_null_stream = self.compute_uncalibrated_frequency_domain_null_stream()
|
|
147
|
+
# Transform to time-frequency domain
|
|
148
|
+
uncalibrated_time_frequency_domain_null_stream = np.array(
|
|
149
|
+
[
|
|
150
|
+
self.time_frequency_transform.frequency_to_wavelet(frequency_domain_data=data)
|
|
151
|
+
for data in uncalibrated_frequency_domain_null_stream
|
|
152
|
+
]
|
|
153
|
+
)
|
|
154
|
+
uncalibrated_time_frequency_domain_null_stream[:, ~self.time_frequency_filter] = 0.0
|
|
155
|
+
return uncalibrated_time_frequency_domain_null_stream
|
|
156
|
+
|
|
157
|
+
def compute_calibrated_time_frequency_domain_null_stream(self, calibration_factor: np.ndarray) -> np.ndarray:
|
|
158
|
+
"""Compute the calibrated time-frequency domain null stream, confined to the filter.
|
|
159
|
+
|
|
160
|
+
Like its uncalibrated counterpart, the returned array is zero outside
|
|
161
|
+
``time_frequency_filter``. The two are deliberately symmetric: a caller that used one
|
|
162
|
+
filtered method and one unfiltered one would sum energies over different pixel domains and
|
|
163
|
+
recreate exactly the normalisation defect this pair was fixed for, silently and without an
|
|
164
|
+
exception.
|
|
165
|
+
|
|
166
|
+
Args:
|
|
167
|
+
calibration_factor (np.ndarray): Calibration factor.
|
|
168
|
+
|
|
169
|
+
Returns:
|
|
170
|
+
np.ndarray: Calibrated time-frequency domain null stream, zero outside the filter.
|
|
171
|
+
"""
|
|
172
|
+
calibrated_frequency_domain_null_stream = self.compute_calibrated_frequency_domain_null_stream(
|
|
173
|
+
calibration_factor=calibration_factor
|
|
174
|
+
)
|
|
175
|
+
# Transform to time-frequency domain
|
|
176
|
+
calibrated_time_frequency_domain_null_stream = np.array(
|
|
177
|
+
[
|
|
178
|
+
self.time_frequency_transform.frequency_to_wavelet(frequency_domain_data=data)
|
|
179
|
+
for data in calibrated_frequency_domain_null_stream
|
|
180
|
+
]
|
|
181
|
+
)
|
|
182
|
+
calibrated_time_frequency_domain_null_stream[:, ~self.time_frequency_filter] = 0.0
|
|
183
|
+
return calibrated_time_frequency_domain_null_stream
|
|
184
|
+
|
|
185
|
+
def construct_calibration_factor_from_parameters(self, parameters: dict) -> np.ndarray:
|
|
186
|
+
"""Construct the calibration factor from parameters.
|
|
187
|
+
|
|
188
|
+
Args:
|
|
189
|
+
parameters (dict): Calibration parameters.
|
|
190
|
+
|
|
191
|
+
Returns:
|
|
192
|
+
np.ndarray: Calibration factor.
|
|
193
|
+
"""
|
|
194
|
+
detector_nodes = []
|
|
195
|
+
expected_keys = set()
|
|
196
|
+
for name in self.interferometers.name:
|
|
197
|
+
prefix = f"recalib_{name}_"
|
|
198
|
+
frequency_indices = sorted(
|
|
199
|
+
int(key.removeprefix(f"{prefix}frequency_"))
|
|
200
|
+
for key in parameters
|
|
201
|
+
if key.startswith(f"{prefix}frequency_") and key.removeprefix(f"{prefix}frequency_").isdigit()
|
|
202
|
+
)
|
|
203
|
+
if (
|
|
204
|
+
frequency_indices != list(range(len(frequency_indices)))
|
|
205
|
+
or len(frequency_indices) < MINIMUM_CUBIC_SPLINE_KNOTS
|
|
206
|
+
):
|
|
207
|
+
raise ValueError("calibration parameters do not match the required detector spline nodes")
|
|
208
|
+
keys = {
|
|
209
|
+
f"{prefix}{quantity}_{index}"
|
|
210
|
+
for quantity in ("frequency", "amplitude", "phase")
|
|
211
|
+
for index in frequency_indices
|
|
212
|
+
}
|
|
213
|
+
expected_keys.update(keys)
|
|
214
|
+
detector_nodes.append((prefix, frequency_indices))
|
|
215
|
+
|
|
216
|
+
supplied_keys = {key for key in parameters if key.startswith("recalib_")}
|
|
217
|
+
if supplied_keys != expected_keys:
|
|
218
|
+
missing = sorted(expected_keys - supplied_keys)
|
|
219
|
+
unexpected = sorted(supplied_keys - expected_keys)
|
|
220
|
+
raise ValueError(
|
|
221
|
+
f"calibration parameters do not match the required keys; missing={missing}, unexpected={unexpected}"
|
|
222
|
+
)
|
|
223
|
+
|
|
224
|
+
calibration_factor_array = np.asarray(
|
|
225
|
+
[
|
|
226
|
+
calibration_factor(
|
|
227
|
+
self.masked_frequency_array,
|
|
228
|
+
[parameters[f"{prefix}frequency_{index}"] for index in indices],
|
|
229
|
+
[parameters[f"{prefix}amplitude_{index}"] for index in indices],
|
|
230
|
+
[parameters[f"{prefix}phase_{index}"] for index in indices],
|
|
231
|
+
)
|
|
232
|
+
for prefix, indices in detector_nodes
|
|
233
|
+
]
|
|
234
|
+
)
|
|
235
|
+
output = np.zeros_like(self._whitened_frequency_domain_strain_array)
|
|
236
|
+
output[:, self.frequency_mask] = calibration_factor_array
|
|
237
|
+
|
|
238
|
+
return output
|
|
239
|
+
|
|
240
|
+
def compute_calibrated_time_frequency_domain_null_stream_from_parameters(self, parameters: dict) -> np.ndarray:
|
|
241
|
+
"""Compute the calibrated time-frequency domain null stream from parameters.
|
|
242
|
+
|
|
243
|
+
Args:
|
|
244
|
+
parameters (dict): A dictionary of calibration parameters.
|
|
245
|
+
|
|
246
|
+
Returns:
|
|
247
|
+
np.ndarray: Calibrated time-frequency domain null stream.
|
|
248
|
+
"""
|
|
249
|
+
calibration_factor = self.construct_calibration_factor_from_parameters(parameters)
|
|
250
|
+
# The filtering lives in compute_calibrated_time_frequency_domain_null_stream, which
|
|
251
|
+
# guarantees confinement for every caller rather than only for this wrapper.
|
|
252
|
+
return self.compute_calibrated_time_frequency_domain_null_stream(calibration_factor=calibration_factor)
|
|
@@ -0,0 +1,48 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Functions for projections.
|
|
3
|
+
"""
|
|
4
|
+
|
|
5
|
+
from __future__ import annotations
|
|
6
|
+
|
|
7
|
+
import numpy as np
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
def compute_projector(
|
|
11
|
+
calibrated_whitened_antenna_response_function: np.ndarray, frequency_mask: np.ndarray
|
|
12
|
+
) -> np.ndarray:
|
|
13
|
+
"""Compute the projector.
|
|
14
|
+
|
|
15
|
+
Args:
|
|
16
|
+
calibrated_whitened_antenna_response_function (np.ndarray): Calibrated whitened antenna response function.
|
|
17
|
+
Dimensions: (frequency, detector, mode).
|
|
18
|
+
frequency_mask (np.ndarray): Frequency mask.
|
|
19
|
+
|
|
20
|
+
Raises:
|
|
21
|
+
ValueError: Frequency mask shape mismatch. The frequency dimensions must have the same length.
|
|
22
|
+
|
|
23
|
+
Returns:
|
|
24
|
+
np.ndarray: Projector. Dimensions: (frequency, detector, detector).
|
|
25
|
+
"""
|
|
26
|
+
n_freq, n_det, _ = calibrated_whitened_antenna_response_function.shape
|
|
27
|
+
if n_freq != frequency_mask.shape[0]:
|
|
28
|
+
raise ValueError(
|
|
29
|
+
"Frequency mask shape mismatch."
|
|
30
|
+
"calibrated_whitened_antenna_response_function:"
|
|
31
|
+
f"(frequency={n_freq},detector={n_det},mode={_})"
|
|
32
|
+
f"and frequency_mask: (frequency={frequency_mask.shape[0]})."
|
|
33
|
+
"The frequency dimensions must have the same length."
|
|
34
|
+
)
|
|
35
|
+
calibrated_whitened_antenna_response_function_masked = calibrated_whitened_antenna_response_function[
|
|
36
|
+
frequency_mask, :, :
|
|
37
|
+
]
|
|
38
|
+
output = np.eye(n_det, dtype=calibrated_whitened_antenna_response_function.dtype)[np.newaxis, :, :].repeat(
|
|
39
|
+
n_freq, axis=0
|
|
40
|
+
)
|
|
41
|
+
f_dagger = np.transpose(np.conj(calibrated_whitened_antenna_response_function_masked), axes=(0, 2, 1))
|
|
42
|
+
# I - F@(F'F)^-1 F'
|
|
43
|
+
projector_gw = np.matmul(
|
|
44
|
+
calibrated_whitened_antenna_response_function_masked,
|
|
45
|
+
np.matmul(np.linalg.inv(np.matmul(f_dagger, calibrated_whitened_antenna_response_function_masked)), f_dagger),
|
|
46
|
+
)
|
|
47
|
+
output[frequency_mask, :, :] = output[frequency_mask, :, :] - projector_gw
|
|
48
|
+
return output
|