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,106 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Functions for whitening.
|
|
3
|
+
"""
|
|
4
|
+
|
|
5
|
+
from __future__ import annotations
|
|
6
|
+
|
|
7
|
+
import numpy as np
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
def compute_whitened_frequency_domain_strain(
|
|
11
|
+
frequency_domain_strain_array: np.ndarray,
|
|
12
|
+
power_spectral_density_array: np.ndarray,
|
|
13
|
+
delta_f: float,
|
|
14
|
+
frequency_mask: np.ndarray,
|
|
15
|
+
) -> np.ndarray:
|
|
16
|
+
"""Compute the whitened frequency domain strain.
|
|
17
|
+
|
|
18
|
+
Args:
|
|
19
|
+
frequency_domain_strain_array (np.ndarray): Frequency domain strain array.
|
|
20
|
+
Dimensions: (detector, frequency).
|
|
21
|
+
power_spectral_density_array (np.ndarray): Power spectral density array.
|
|
22
|
+
Dimensions: (detector, frequency).
|
|
23
|
+
delta_f (float): Frequency resolution in Hz.
|
|
24
|
+
frequency_mask (np.ndarray): A frequency mask.
|
|
25
|
+
|
|
26
|
+
Raises:
|
|
27
|
+
ValueError: delta_f must be positive.
|
|
28
|
+
ValueError: Shape mismatch. frequency_domain_strain_array and power_spectral_density_array
|
|
29
|
+
must have the same dimensions.
|
|
30
|
+
ValueError: Shape mismatch. frequency_domain_strain_array and frequency_mask
|
|
31
|
+
must have the same frequency dimension.
|
|
32
|
+
|
|
33
|
+
Returns:
|
|
34
|
+
np.ndarray: Whitened frequency domain strain.
|
|
35
|
+
"""
|
|
36
|
+
if delta_f <= 0.0:
|
|
37
|
+
raise ValueError("delta_f must be positive.")
|
|
38
|
+
if frequency_domain_strain_array.shape != power_spectral_density_array.shape:
|
|
39
|
+
raise ValueError(
|
|
40
|
+
"Shape mismatch. frequency_domain_strain_array:"
|
|
41
|
+
f"{frequency_domain_strain_array.shape}"
|
|
42
|
+
f"and power_spectral_density_array: {power_spectral_density_array.shape}"
|
|
43
|
+
"must have the same dimensions."
|
|
44
|
+
)
|
|
45
|
+
if frequency_domain_strain_array.shape[1] != frequency_mask.shape[0]:
|
|
46
|
+
raise ValueError(
|
|
47
|
+
"Shape mismatch. frequency_domain_strain_array:"
|
|
48
|
+
f"{frequency_domain_strain_array.shape} and frequency_mask:"
|
|
49
|
+
f"{frequency_mask.shape} must have the same frequency dimension."
|
|
50
|
+
)
|
|
51
|
+
return np.divide(
|
|
52
|
+
frequency_domain_strain_array,
|
|
53
|
+
np.sqrt(power_spectral_density_array / (2 * delta_f)),
|
|
54
|
+
where=(frequency_mask & np.all(power_spectral_density_array, axis=0)),
|
|
55
|
+
out=np.zeros_like(frequency_domain_strain_array),
|
|
56
|
+
)
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
def compute_whitened_antenna_response(
|
|
60
|
+
antenna_response_matrix: np.ndarray,
|
|
61
|
+
power_spectral_density_array: np.ndarray,
|
|
62
|
+
delta_f: float,
|
|
63
|
+
frequency_mask: np.ndarray,
|
|
64
|
+
) -> np.ndarray:
|
|
65
|
+
"""Compute the whitened antenna response function.
|
|
66
|
+
|
|
67
|
+
Args:
|
|
68
|
+
antenna_response_matrix (np.ndarray): Antenna response matrix.
|
|
69
|
+
power_spectral_density_array (np.ndarray): Power spectral density array.
|
|
70
|
+
delta_f (float): Frequency resolution in Hz.
|
|
71
|
+
frequency_mask (np.ndarray): Frequency mask.
|
|
72
|
+
|
|
73
|
+
Raises:
|
|
74
|
+
ValueError: delta_f must be positive.
|
|
75
|
+
ValueError: Shape mismatch. antenna_response_matrix and power_spectral_density_array
|
|
76
|
+
must have the same detector dimension.
|
|
77
|
+
ValueError: Shape mismatch. power_spectral_density_array and frequency_mask
|
|
78
|
+
must have the same frequency dimensions.
|
|
79
|
+
|
|
80
|
+
Returns:
|
|
81
|
+
np.ndarray: Whitened antenna response function.
|
|
82
|
+
"""
|
|
83
|
+
if delta_f <= 0:
|
|
84
|
+
raise ValueError("delta_f must be positive.")
|
|
85
|
+
n_det_1, n_mode_1 = antenna_response_matrix.shape
|
|
86
|
+
n_det_2, n_freq_2 = power_spectral_density_array.shape
|
|
87
|
+
n_freq_3 = frequency_mask.shape[0]
|
|
88
|
+
if n_det_1 != n_det_2:
|
|
89
|
+
raise ValueError(
|
|
90
|
+
f"Shape mismatch. antenna_response_matrix: (detector={n_det_1}, mode={n_mode_1})"
|
|
91
|
+
f"and power_spectral_density_array: (detector={n_det_2}, frequency={n_freq_2})"
|
|
92
|
+
"must have the same detector dimension."
|
|
93
|
+
)
|
|
94
|
+
if n_freq_2 != n_freq_3:
|
|
95
|
+
raise ValueError(
|
|
96
|
+
f"Shape mismatch. power_spectral_density_array: (detector={n_det_2}, freq={n_freq_2})"
|
|
97
|
+
f"and frequency_mask: (frequency: {n_freq_3}) must have the same frequency dimension."
|
|
98
|
+
)
|
|
99
|
+
output = np.zeros((n_freq_2, n_det_1, n_mode_1), dtype=antenna_response_matrix.dtype)
|
|
100
|
+
frequency_mask = frequency_mask & np.all(power_spectral_density_array, axis=0)
|
|
101
|
+
output[frequency_mask, :, :] = np.einsum(
|
|
102
|
+
"dm,df->fdm",
|
|
103
|
+
antenna_response_matrix,
|
|
104
|
+
1 / np.sqrt(power_spectral_density_array[:, frequency_mask] / (2 * delta_f)),
|
|
105
|
+
)
|
|
106
|
+
return output
|
nullcal/result/result.py
ADDED
|
@@ -0,0 +1,50 @@
|
|
|
1
|
+
"""Sampler result types."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from collections.abc import Mapping
|
|
6
|
+
from dataclasses import dataclass, field
|
|
7
|
+
from typing import Any
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
@dataclass(frozen=True)
|
|
11
|
+
class Result:
|
|
12
|
+
"""Posterior samples extracted from a sequence of sampler states."""
|
|
13
|
+
|
|
14
|
+
samples: Mapping[str, Any]
|
|
15
|
+
logdensity: Any
|
|
16
|
+
info: Any | None = None
|
|
17
|
+
metadata: Mapping[str, Any] = field(default_factory=dict)
|
|
18
|
+
|
|
19
|
+
def __post_init__(self) -> None:
|
|
20
|
+
if not self.samples:
|
|
21
|
+
raise ValueError("samples must contain at least one parameter array")
|
|
22
|
+
sample_counts = {value.shape[0] for value in self.samples.values() if value.ndim > 0}
|
|
23
|
+
if len(sample_counts) != 1 or any(value.ndim == 0 for value in self.samples.values()):
|
|
24
|
+
raise ValueError("all parameter arrays must have the same leading sample dimension")
|
|
25
|
+
if self.logdensity.ndim != 1 or self.logdensity.shape[0] != self.sample_count:
|
|
26
|
+
raise ValueError("logdensity must contain one value per posterior sample")
|
|
27
|
+
|
|
28
|
+
@property
|
|
29
|
+
def sample_count(self) -> int:
|
|
30
|
+
"""Number of retained posterior samples."""
|
|
31
|
+
return next(iter(self.samples.values())).shape[0]
|
|
32
|
+
|
|
33
|
+
@classmethod
|
|
34
|
+
def from_states(
|
|
35
|
+
cls,
|
|
36
|
+
states: Any,
|
|
37
|
+
*,
|
|
38
|
+
info: Any | None = None,
|
|
39
|
+
metadata: Mapping[str, Any] | None = None,
|
|
40
|
+
) -> Result:
|
|
41
|
+
"""Extract positions and log densities from stacked sampler states."""
|
|
42
|
+
return cls(
|
|
43
|
+
samples=states.position,
|
|
44
|
+
logdensity=states.logdensity,
|
|
45
|
+
info=info,
|
|
46
|
+
metadata={} if metadata is None else metadata,
|
|
47
|
+
)
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
__all__ = ["Result"]
|
nullcal/result/utils.py
ADDED
|
@@ -0,0 +1,13 @@
|
|
|
1
|
+
"""Small result-processing helpers."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
def spline_percentage_xform(delta_a: float | np.ndarray) -> float | np.ndarray:
|
|
9
|
+
"""Convert fractional amplitude error to percent."""
|
|
10
|
+
return delta_a * 100
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
__all__ = ["spline_percentage_xform"]
|
nullcal/sampler.py
ADDED
|
@@ -0,0 +1,113 @@
|
|
|
1
|
+
"""BlackJAX NUTS sampling with convergence diagnostics."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from collections.abc import Callable, Mapping
|
|
6
|
+
from typing import Any
|
|
7
|
+
|
|
8
|
+
import blackjax
|
|
9
|
+
import jax
|
|
10
|
+
import jax.numpy as jnp
|
|
11
|
+
import numpy as np
|
|
12
|
+
|
|
13
|
+
from .result import Result
|
|
14
|
+
|
|
15
|
+
MINIMUM_DIAGNOSTIC_CHAINS = 2
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
def _jitter_position(position, key, scale: float):
|
|
19
|
+
leaves, structure = jax.tree.flatten(position)
|
|
20
|
+
keys = jax.random.split(key, len(leaves))
|
|
21
|
+
jittered = [
|
|
22
|
+
jnp.asarray(leaf) + scale * jax.random.normal(leaf_key, shape=jnp.asarray(leaf).shape)
|
|
23
|
+
for leaf, leaf_key in zip(leaves, keys, strict=True)
|
|
24
|
+
]
|
|
25
|
+
return jax.tree.unflatten(structure, jittered)
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
def sample_nuts(
|
|
29
|
+
logdensity_fn: Callable[[Any], jax.Array],
|
|
30
|
+
initial_position: Mapping[str, Any],
|
|
31
|
+
*,
|
|
32
|
+
seed: int = 0,
|
|
33
|
+
num_chains: int = 4,
|
|
34
|
+
num_warmup: int = 1_000,
|
|
35
|
+
num_samples: int = 1_000,
|
|
36
|
+
target_acceptance_rate: float = 0.8,
|
|
37
|
+
initial_position_jitter: float = 0.01,
|
|
38
|
+
) -> Result:
|
|
39
|
+
"""Adapt and run independent NUTS chains, returning mandatory diagnostics.
|
|
40
|
+
|
|
41
|
+
Chains are executed independently rather than treating detector data as a
|
|
42
|
+
realization batch. Returned parameter arrays flatten chain and draw into a
|
|
43
|
+
single leading sample dimension; chain-aware R-hat and ESS are retained in
|
|
44
|
+
``Result.metadata``.
|
|
45
|
+
"""
|
|
46
|
+
if num_chains < MINIMUM_DIAGNOSTIC_CHAINS:
|
|
47
|
+
raise ValueError("at least two chains are required to compute R-hat")
|
|
48
|
+
if num_warmup <= 0 or num_samples <= 0:
|
|
49
|
+
raise ValueError("num_warmup and num_samples must be positive")
|
|
50
|
+
if initial_position_jitter < 0.0:
|
|
51
|
+
raise ValueError("initial_position_jitter must be non-negative")
|
|
52
|
+
|
|
53
|
+
root_key = jax.random.key(seed)
|
|
54
|
+
chain_keys = jax.random.split(root_key, num_chains)
|
|
55
|
+
chain_states = []
|
|
56
|
+
chain_infos = []
|
|
57
|
+
|
|
58
|
+
adaptation = blackjax.window_adaptation(
|
|
59
|
+
blackjax.nuts,
|
|
60
|
+
logdensity_fn,
|
|
61
|
+
target_acceptance_rate=target_acceptance_rate,
|
|
62
|
+
)
|
|
63
|
+
|
|
64
|
+
for chain_key in chain_keys:
|
|
65
|
+
jitter_key, warmup_key, sample_key = jax.random.split(chain_key, 3)
|
|
66
|
+
chain_initial_position = _jitter_position(initial_position, jitter_key, initial_position_jitter)
|
|
67
|
+
adapted, _ = adaptation.run(warmup_key, chain_initial_position, num_steps=num_warmup)
|
|
68
|
+
kernel = blackjax.nuts(logdensity_fn, **adapted.parameters)
|
|
69
|
+
sample_keys = jax.random.split(sample_key, num_samples)
|
|
70
|
+
|
|
71
|
+
def step(state, key, kernel=kernel):
|
|
72
|
+
next_state, info = kernel.step(key, state)
|
|
73
|
+
return next_state, (next_state, info)
|
|
74
|
+
|
|
75
|
+
_, (states, infos) = jax.lax.scan(step, adapted.state, sample_keys)
|
|
76
|
+
chain_states.append(states)
|
|
77
|
+
chain_infos.append(infos)
|
|
78
|
+
|
|
79
|
+
positions_by_chain = jax.tree.map(lambda *values: jnp.stack(values), *(state.position for state in chain_states))
|
|
80
|
+
logdensity_by_chain = jnp.stack([state.logdensity for state in chain_states])
|
|
81
|
+
divergence_by_chain = jnp.stack([info.is_divergent for info in chain_infos])
|
|
82
|
+
|
|
83
|
+
diagnostic_leaves = [value.reshape((num_chains, num_samples, -1)) for value in jax.tree.leaves(positions_by_chain)]
|
|
84
|
+
diagnostic_array = jnp.concatenate(diagnostic_leaves, axis=-1)
|
|
85
|
+
rhat = blackjax.diagnostics.rhat(diagnostic_array)
|
|
86
|
+
ess_bulk = blackjax.diagnostics.ess_bulk(diagnostic_array)
|
|
87
|
+
ess_tail = blackjax.diagnostics.ess_tail(diagnostic_array)
|
|
88
|
+
|
|
89
|
+
flattened_positions = jax.tree.map(
|
|
90
|
+
lambda value: value.reshape((num_chains * num_samples, *value.shape[2:])), positions_by_chain
|
|
91
|
+
)
|
|
92
|
+
flattened_logdensity = logdensity_by_chain.reshape(num_chains * num_samples)
|
|
93
|
+
flattened_divergence = divergence_by_chain.reshape(num_chains * num_samples)
|
|
94
|
+
metadata = {
|
|
95
|
+
"seed": seed,
|
|
96
|
+
"num_chains": num_chains,
|
|
97
|
+
"num_warmup": num_warmup,
|
|
98
|
+
"num_samples_per_chain": num_samples,
|
|
99
|
+
"target_acceptance_rate": target_acceptance_rate,
|
|
100
|
+
"max_rhat": float(np.asarray(jnp.max(rhat))),
|
|
101
|
+
"min_ess_bulk": float(np.asarray(jnp.min(ess_bulk))),
|
|
102
|
+
"min_ess_tail": float(np.asarray(jnp.min(ess_tail))),
|
|
103
|
+
"divergences": int(np.asarray(jnp.sum(divergence_by_chain))),
|
|
104
|
+
}
|
|
105
|
+
return Result(
|
|
106
|
+
samples=flattened_positions,
|
|
107
|
+
logdensity=flattened_logdensity,
|
|
108
|
+
info={"is_divergent": flattened_divergence},
|
|
109
|
+
metadata=metadata,
|
|
110
|
+
)
|
|
111
|
+
|
|
112
|
+
|
|
113
|
+
__all__ = ["sample_nuts"]
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
"""Numerical studies that inform nullcal model choices."""
|