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.
Files changed (42) hide show
  1. nullcal/__init__.py +16 -0
  2. nullcal/calibration.py +189 -0
  3. nullcal/clustering/__init__.py +0 -0
  4. nullcal/clustering/base.py +39 -0
  5. nullcal/clustering/injection.py +63 -0
  6. nullcal/clustering/precompute.py +30 -0
  7. nullcal/clustering/single.py +208 -0
  8. nullcal/clustering/time_frequency_map.py +42 -0
  9. nullcal/data.py +93 -0
  10. nullcal/likelihood/__init__.py +9 -0
  11. nullcal/likelihood/recalibration_likelihood.py +227 -0
  12. nullcal/metadata/__init__.py +0 -0
  13. nullcal/metadata/yaml.py +30 -0
  14. nullcal/null_stream/__init__.py +0 -0
  15. nullcal/null_stream/calibration.py +40 -0
  16. nullcal/null_stream/null_stream.py +252 -0
  17. nullcal/null_stream/projector.py +48 -0
  18. nullcal/null_stream/whiten.py +106 -0
  19. nullcal/result/__init__.py +5 -0
  20. nullcal/result/result.py +50 -0
  21. nullcal/result/utils.py +13 -0
  22. nullcal/sampler.py +113 -0
  23. nullcal/studies/__init__.py +1 -0
  24. nullcal/studies/lwa_leakage.py +513 -0
  25. nullcal/studies/spline_resolution.py +148 -0
  26. nullcal/time_frequency_transform/README.md +7 -0
  27. nullcal/time_frequency_transform/__init__.py +23 -0
  28. nullcal/time_frequency_transform/inverse_wavelet_freq_funcs.py +42 -0
  29. nullcal/time_frequency_transform/inverse_wavelet_time_funcs.py +49 -0
  30. nullcal/time_frequency_transform/stft.py +45 -0
  31. nullcal/time_frequency_transform/transform_freq_funcs.py +180 -0
  32. nullcal/time_frequency_transform/transform_time_funcs.py +60 -0
  33. nullcal/time_frequency_transform/utils.py +21 -0
  34. nullcal/time_frequency_transform/wavelet_transforms.py +249 -0
  35. nullcal/utils/__init__.py +0 -0
  36. nullcal/utils/log.py +72 -0
  37. nullcal/utils/snr.py +29 -0
  38. nullcal/version.py +9 -0
  39. nullcal-0.2.0.dist-info/METADATA +222 -0
  40. nullcal-0.2.0.dist-info/RECORD +42 -0
  41. nullcal-0.2.0.dist-info/WHEEL +4 -0
  42. nullcal-0.2.0.dist-info/licenses/LICENSE +21 -0
@@ -0,0 +1,9 @@
1
+ """A submodule for likelihoods."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from .recalibration_likelihood import RecalibrationLikelihood
6
+
7
+ __all__ = [
8
+ "RecalibrationLikelihood",
9
+ ]
@@ -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
@@ -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