From b65a70c8b9c8d361497d86d0d8c1717122c69016 Mon Sep 17 00:00:00 2001 From: yallup Date: Thu, 16 Jul 2026 12:29:06 +0100 Subject: [PATCH 1/2] Add cache-aware Nested SwiG sampler --- docs/guides/cli.md | 16 ++ docs/guides/samplers.md | 56 +++- examples/GW170817_SwiG.py | 102 ++++++++ src/jimgw/core/jim.py | 170 ++++++++++++- src/jimgw/core/single_event/likelihood.py | 98 ++++++- src/jimgw/samplers/__init__.py | 22 +- src/jimgw/samplers/blackjax/nss.py | 40 ++- src/jimgw/samplers/blackjax/swig.py | 239 ++++++++++++++++++ src/jimgw/samplers/config.py | 74 +++++- .../unit/core/single_event/test_likelihood.py | 204 +++++++++++++++ tests/unit/samplers/blackjax/test_swig.py | 62 +++++ tests/unit/samplers/test_config.py | 22 ++ tests/unit/samplers/test_registry.py | 3 +- 13 files changed, 1078 insertions(+), 30 deletions(-) create mode 100644 examples/GW170817_SwiG.py create mode 100644 src/jimgw/samplers/blackjax/swig.py create mode 100644 tests/unit/samplers/blackjax/test_swig.py diff --git a/docs/guides/cli.md b/docs/guides/cli.md index f3ac1d38..4d50b993 100644 --- a/docs/guides/cli.md +++ b/docs/guides/cli.md @@ -300,6 +300,7 @@ The `type` field selects the backend. Each backend has its own set of tuning par | `flowmc` | Normalizing-flow MCMC | No | | `blackjax-smc` | Sequential Monte Carlo | Yes | | `blackjax-nss` | Nested slice sampling | Yes | +| `blackjax-swig` | Nested Slice within Gibbs with waveform caching | Yes | | `blackjax-ns-aw` | Nested sampling (acceptance-walk) | Yes | ### `type = "flowmc"` @@ -356,6 +357,21 @@ n_tempered_steps = 5 | `checkpoint_dir` | `{output.dir}/` | Directory for `checkpoint.pkl`; set by the CLI automatically | | `checkpoint_interval` | `600.0` | Seconds between checkpoint writes; `0` disables checkpointing | +### `type = "blackjax-swig"` + +| Field | Default | Description | +| --- | --- | --- | +| `blocks` | required | Ordered lists of sampling-space parameter names | +| `n_live` | `512` | Number of live points | +| `n_delete_frac` | `0.125` | Fraction of live points replaced per iteration | +| `num_inner_steps_per_dim` | `1` | Slice steps per dimension within each block | +| `num_gibbs_sweeps` | `2` | Complete block sweeps per replacement | +| `max_steps` | `10` | Maximum stepping-out expansions per slice | +| `max_shrinkage` | `100` | Maximum shrinkage evaluations per slice | +| `termination_dlogz` | `exp(-3)` (~`0.0498`) | Stop when the remaining evidence contribution falls below this | +| `checkpoint_dir` | `{output.dir}/` | Directory for `checkpoint.pkl`; set by the CLI automatically | +| `checkpoint_interval` | `600.0` | Seconds between checkpoint writes; `0` disables checkpointing | + ### `type = "blackjax-ns-aw"` Requires all sampling-space parameters to lie in $[0, 1]$. The CLI enforces this automatically. diff --git a/docs/guides/samplers.md b/docs/guides/samplers.md index df03d5cb..206fda07 100644 --- a/docs/guides/samplers.md +++ b/docs/guides/samplers.md @@ -16,6 +16,7 @@ samples = jim.get_samples() # dict[str, np.ndarray] keyed by parameter name | [flowMC](#flowmc) | normalizing-flow-enhanced MCMC | No | None | | [NS-AW](#blackjax-ns-aw) | Nested sampling (bilby/dynesty-style acceptance-walk) | Yes | Uniform prior; unit-cube sampling space | | [NSS](#blackjax-nss) | Nested slice sampling | Yes | Normalised prior | +| [SwiG](#blackjax-swig) | Nested Slice within Gibbs with waveform caching | Yes | Normalised prior | | [SMC](#blackjax-smc) | Sequential Monte Carlo | Yes | Normalised prior | --- @@ -168,6 +169,58 @@ Key parameters: --- +### BlackJAX SwiG + +Nested Slice within Gibbs (SwiG) applies covariance-shaped slice updates to +named parameter blocks. With `TransientLikelihoodFD` or +`HeterodynedTransientLikelihoodFD`, Jim caches the generated waveform +polarizations: blocks that affect waveform inputs regenerate the cache, while +detector-projection-only blocks reuse it. Luminosity distance is also factored +out analytically, so a `d_L` block reuses the cache. Jim determines the other +dependencies after sample transforms, likelihood transforms, fixed parameters, +and analytic marginalisation have been applied. + +```python +import math + +from jimgw.samplers.config import BlackJAXSwiGConfig + +jim = Jim( + likelihood, + prior, + sampler_config=BlackJAXSwiGConfig( + blocks=[ + ["M_c", "q", "lambda_1", "lambda_2"], + ["s1_z", "s2_z"], + ["iota"], + ["ra", "dec"], + ["psi"], + ["t_c"], + ], + n_live=512, + n_delete_frac=0.125, + num_inner_steps_per_dim=1, + num_gibbs_sweeps=2, + termination_dlogz=math.exp(-3.0), + ), + likelihood_transforms=likelihood_transforms, +) +``` + +Blocks must form an exact partition of the sampling parameters. A block that +contains any waveform dependency is treated as cache-refreshing; mixed blocks +are therefore safe but may be less efficient. Parameters that are analytically +marginalised must not appear in the blocks. + +Key parameters: + +- `blocks` — ordered sampling-space parameter blocks for each Gibbs sweep. +- `n_live` / `n_delete_frac` — live-set size and replacement fraction, matching NSS. +- `num_inner_steps_per_dim` — random-direction slice steps per block dimension. +- `num_gibbs_sweeps` — complete block sweeps per particle replacement. + +--- + ## Checkpointing and resuming All samplers support checkpoint/resume so long-running jobs can survive interruptions. @@ -205,7 +258,8 @@ jim = Jim( jim.sample() # resumes from ./my_run/checkpoint.pkl ``` -The same fields work identically for `FlowMCConfig`, `BlackJAXNSAWConfig`, and `BlackJAXNSSConfig`. +The same fields work identically for `FlowMCConfig`, `BlackJAXNSAWConfig`, +`BlackJAXNSSConfig`, and `BlackJAXSwiGConfig`. | Field | Default | Notes | | --- | --- | --- | diff --git a/examples/GW170817_SwiG.py b/examples/GW170817_SwiG.py new file mode 100644 index 00000000..91ef7ed7 --- /dev/null +++ b/examples/GW170817_SwiG.py @@ -0,0 +1,102 @@ +"""GW170817-style 128 s BNS analysis with BlackJAX Nested SwiG. + +This example deliberately uses Ripple's built-in aligned-spin +IMRPhenomD_NRTidalv2 model. Coalescence time and phase are analytically +marginalized; the remaining parameters are partitioned into waveform and +projection blocks. +""" + +import time +import math + +import jax +import jax.numpy as jnp +from jimgw.core.jim import Jim +from jimgw.core.prior import ( + CombinePrior, + CosinePrior, + PowerLawPrior, + SinePrior, + UniformPrior, +) +from jimgw.core.single_event.data import Data +from jimgw.core.single_event.detector import get_H1, get_L1, get_V1 +from jimgw.core.single_event.likelihood import TransientLikelihoodFD +from jimgw.core.single_event.transforms import MassRatioToSymmetricMassRatioTransform +from jimgw.core.single_event.waveform import RippleIMRPhenomD_NRTidalv2 +from jimgw.samplers.config import BlackJAXSwiGConfig + +jax.config.update("jax_enable_x64", True) + + +gps = 1187008882.43 +duration = 128.0 +start = gps + 2.0 - duration +end = start + duration +psd_start = start - 2048.0 + +ifos = [get_H1(), get_L1(), get_V1()] +for ifo in ifos: + strain = Data.from_gwosc(ifo.name, start, end) + ifo.set_data(strain) + psd_data = Data.from_gwosc(ifo.name, psd_start, start) + ifo.set_psd( + psd_data.to_psd(nperseg=int(strain.duration * strain.sampling_frequency)) + ) + +distance_prior = PowerLawPrior(30.0, 150.0, 2.0, parameter_names=["d_L"]) +prior = CombinePrior( + [ + UniformPrior(1.18, 1.21, parameter_names=["M_c"]), + UniformPrior(0.5, 1.0, parameter_names=["q"]), + UniformPrior(-0.05, 0.05, parameter_names=["s1_z"]), + UniformPrior(-0.05, 0.05, parameter_names=["s2_z"]), + UniformPrior(0.0, 5000.0, parameter_names=["lambda_1"]), + UniformPrior(0.0, 5000.0, parameter_names=["lambda_2"]), + SinePrior(parameter_names=["iota"]), + distance_prior, + UniformPrior(0.0, 2.0 * jnp.pi, parameter_names=["ra"]), + CosinePrior(parameter_names=["dec"]), + UniformPrior(0.0, jnp.pi, parameter_names=["psi"]), + ] +) + +likelihood = TransientLikelihoodFD( + ifos, + waveform=RippleIMRPhenomD_NRTidalv2(f_ref=20.0), + trigger_time=gps, + f_min=20.0, + f_max=2048.0, + phase_marginalization=True, + time_marginalization={"tc_range": (-0.03, 0.03)}, +) + +jim = Jim( + likelihood, + prior, + likelihood_transforms=[MassRatioToSymmetricMassRatioTransform], + periodic={ + "ra": (0.0, 2.0 * float(jnp.pi)), + "psi": (0.0, float(jnp.pi)), + }, + sampler_config=BlackJAXSwiGConfig( + blocks=[ + ["M_c", "q", "lambda_1", "lambda_2"], + ["s1_z", "s2_z"], + ["iota"], + ["d_L"], + ["ra", "dec"], + ["psi"], + ], + n_live=512, + n_delete_frac=0.125, + num_inner_steps_per_dim=1, + num_gibbs_sweeps=2, + termination_dlogz=math.exp(-3.0), + ), +) + +start_time = time.time() +jim.sample() +print(f"Sampling took {(time.time() - start_time) / 60.0:.2f} minutes") +print(jim.get_diagnostics()) diff --git a/src/jimgw/core/jim.py b/src/jimgw/core/jim.py index 47223b13..f23b9c22 100644 --- a/src/jimgw/core/jim.py +++ b/src/jimgw/core/jim.py @@ -12,9 +12,13 @@ from jimgw.core.base import LikelihoodBase from jimgw.core.prior import Prior from jimgw.core.transforms import BijectiveTransform, NtoMTransform -from jimgw.core.single_event.likelihood import SingleEventLikelihood +from jimgw.core.single_event.likelihood import ( + HeterodynedTransientLikelihoodFD, + SingleEventLikelihood, + TransientLikelihoodFD, +) from jimgw.samplers import Sampler, SamplerConfig, build_sampler -from jimgw.samplers.config import FlowMCConfig +from jimgw.samplers.config import BlackJAXSwiGConfig, FlowMCConfig from jimgw._logging import ensure_logger_handler logger = logging.getLogger(__name__) @@ -128,6 +132,24 @@ def __init__( else: periodic_resolved = None + block_indices = None + refresh_cache = None + build_cache_fn = None + log_likelihood_from_cache_fn = None + if isinstance(sampler_config, BlackJAXSwiGConfig): + block_indices, refresh_cache = self._resolve_swig_blocks(sampler_config) + if not isinstance( + likelihood, + (TransientLikelihoodFD, HeterodynedTransientLikelihoodFD), + ): + raise TypeError( + "BlackJAXSwiGConfig requires TransientLikelihoodFD or " + "HeterodynedTransientLikelihoodFD because its cache is the " + "generated waveform polarizations." + ) + build_cache_fn = self._build_cache_fn + log_likelihood_from_cache_fn = self._log_likelihood_from_cache_fn + self.sampler = build_sampler( sampler_config, n_dims=len(self.parameter_names), @@ -135,6 +157,10 @@ def __init__( log_likelihood_fn=self._log_likelihood_fn, log_posterior_fn=self._log_posterior_fn, periodic=periodic_resolved, + block_indices=block_indices, + refresh_cache=refresh_cache, + build_cache_fn=build_cache_fn, + log_likelihood_from_cache_fn=log_likelihood_from_cache_fn, ) self._verify_posterior() @@ -294,6 +320,16 @@ def _setup_problem( # (n_dims,) and are injected into the sampler. names = self.parameter_names + def _to_likelihood_parameters( + arr: Float[Array, " n_dims"], + ) -> dict[str, Float]: + named = dict(zip(names, arr, strict=True)) + for transform in reversed(sample_transforms): + named, _ = transform.inverse(named) + for transform in likelihood_transforms: + named = transform.forward(named) + return named + def _log_prior_fn(arr: Float[Array, " n_dims"]) -> FloatScalar: named = dict(zip(names, arr, strict=True)) jac: FloatScalar = jnp.zeros(()) @@ -303,13 +339,29 @@ def _log_prior_fn(arr: Float[Array, " n_dims"]) -> FloatScalar: return prior.log_prob(named) + jac def _log_likelihood_fn(arr: Float[Array, " n_dims"]) -> FloatScalar: - named = dict(zip(names, arr, strict=True)) - for transform in reversed(sample_transforms): - named, _ = transform.inverse(named) - for transform in likelihood_transforms: - named = transform.forward(named) + named = _to_likelihood_parameters(arr) return likelihood.evaluate(named) + def _build_cache_fn(arr: Float[Array, " n_dims"]): + if not isinstance( + likelihood, + (TransientLikelihoodFD, HeterodynedTransientLikelihoodFD), + ): + raise TypeError("waveform caching requires a transient likelihood") + return likelihood.generate_waveform(_to_likelihood_parameters(arr)) + + def _log_likelihood_from_cache_fn( + arr: Float[Array, " n_dims"], cache + ) -> FloatScalar: + if not isinstance( + likelihood, + (TransientLikelihoodFD, HeterodynedTransientLikelihoodFD), + ): + raise TypeError("waveform caching requires a transient likelihood") + return likelihood.evaluate_from_waveform( + _to_likelihood_parameters(arr), cache + ) + def _log_posterior_fn(arr: Float[Array, " n_dims"]) -> FloatScalar: named = dict(zip(names, arr, strict=True)) jac: FloatScalar = jnp.zeros(()) @@ -324,6 +376,99 @@ def _log_posterior_fn(arr: Float[Array, " n_dims"]) -> FloatScalar: self._log_prior_fn = _log_prior_fn self._log_likelihood_fn = _log_likelihood_fn self._log_posterior_fn = _log_posterior_fn + self._build_cache_fn = _build_cache_fn + self._log_likelihood_from_cache_fn = _log_likelihood_from_cache_fn + + def _resolve_swig_blocks( + self, config: BlackJAXSwiGConfig + ) -> tuple[tuple[tuple[int, ...], ...], tuple[bool, ...]]: + """Validate named SwiG blocks and infer waveform-cache invalidation.""" + names = self.parameter_names + flattened = [name for block in config.blocks for name in block] + marginalized = set() + likelihood = self.likelihood + if isinstance(likelihood, SingleEventLikelihood): + if getattr(likelihood, "time_marginalization", False): + marginalized.add("t_c") + if getattr(likelihood, "phase_marginalization", False): + marginalized.add("phase_c") + if getattr(likelihood, "distance_marginalization", False): + marginalized.add("d_L") + + blocked_marginalized = sorted(marginalized.intersection(flattened)) + if blocked_marginalized: + raise ValueError( + "SwiG blocks contain analytically marginalized parameter(s) " + f"{blocked_marginalized}; marginalized parameters are not sampled." + ) + + unknown = sorted(set(flattened) - set(names)) + missing = sorted(set(names) - set(flattened)) + if unknown: + raise ValueError( + f"SwiG block parameter(s) {unknown} are not sampling parameters {names}." + ) + if missing: + raise ValueError( + f"SwiG blocks do not cover sampling parameter(s) {missing}." + ) + + waveform_dependencies = self._waveform_sampling_dependencies(marginalized) + block_indices = tuple( + tuple(names.index(name) for name in block) for block in config.blocks + ) + refresh_cache = tuple( + bool(set(block).intersection(waveform_dependencies)) + for block in config.blocks + ) + return block_indices, refresh_cache + + def _waveform_sampling_dependencies(self, marginalized: set[str]) -> set[str]: + """Conservatively map waveform inputs back to sampling-space names.""" + all_sampling = set(self.parameter_names) + dependencies: dict[str, set[str]] = { + name: {name} for name in self.parameter_names + } + + def apply_mapping(transform, reverse: bool) -> None: + from_names, to_names = transform.name_mapping + consumed = to_names if reverse else from_names + produced = from_names if reverse else to_names + conditional = getattr(transform, "conditional_names", []) + inputs = list(consumed) + list(conditional) + combined: set[str] = set() + for name in inputs: + combined.update(dependencies.get(name, all_sampling)) + for name in consumed: + dependencies.pop(name, None) + for name in produced: + dependencies[name] = combined.copy() + + for transform in reversed(self.sample_transforms): + apply_mapping(transform, reverse=True) + for transform in self.likelihood_transforms: + apply_mapping(transform, reverse=False) + + likelihood = self.likelihood + if not isinstance( + likelihood, + (TransientLikelihoodFD, HeterodynedTransientLikelihoodFD), + ): + return all_sampling + + waveform_dependencies: set[str] = set() + for parameter in likelihood.waveform.parameter_names: + if parameter == "d_L": + # Cached polarizations factor out inverse-distance amplitude. + continue + if parameter in marginalized: + continue + if parameter in likelihood.fixed_parameters: + if callable(likelihood.fixed_parameters[parameter]): + waveform_dependencies.update(all_sampling) + continue + waveform_dependencies.update(dependencies.get(parameter, all_sampling)) + return waveform_dependencies def _verify_posterior(self) -> None: """Draw test points from the prior and verify the posterior is not mostly NaN. @@ -365,10 +510,17 @@ def _validate_normalized_prior( [`BlackJAXSMCConfig`][jimgw.samplers.config.BlackJAXSMCConfig] and ``prior.is_normalized`` is ``False``. """ - from jimgw.samplers.config import BlackJAXNSSConfig, BlackJAXSMCConfig + from jimgw.samplers.config import ( + BlackJAXNSSConfig, + BlackJAXSMCConfig, + BlackJAXSwiGConfig, + ) if ( - isinstance(sampler_config, (BlackJAXNSSConfig, BlackJAXSMCConfig)) + isinstance( + sampler_config, + (BlackJAXNSSConfig, BlackJAXSwiGConfig, BlackJAXSMCConfig), + ) and not prior.is_normalized ): raise ValueError( diff --git a/src/jimgw/core/single_event/likelihood.py b/src/jimgw/core/single_event/likelihood.py index b6a402ca..8932689c 100644 --- a/src/jimgw/core/single_event/likelihood.py +++ b/src/jimgw/core/single_event/likelihood.py @@ -92,6 +92,20 @@ def __init__( self.waveform = waveform self.fixed_parameters = fixed_parameters if fixed_parameters is not None else {} + def _generate_cached_polarizations(self, frequencies, params): + """Generate polarizations with analytic distance scaling factored out.""" + cache_params = params.copy() + if "d_L" in getattr(self.waveform, "parameter_names", ()): + cache_params["d_L"] = 1.0 + return self.waveform(frequencies, cache_params) + + def _restore_cached_distance(self, waveform_sky, params): + """Apply inverse-distance amplitude scaling to cached polarizations.""" + if "d_L" not in getattr(self.waveform, "parameter_names", ()): + return waveform_sky + scale = 1.0 / params["d_L"] + return {polarization: strain * scale for polarization, strain in waveform_sky.items()} + def evaluate(self, params: dict[str, Float]) -> FloatScalar: """Apply ``fixed_parameters`` overrides and evaluate the likelihood. @@ -270,6 +284,11 @@ def __init__( self._init_distance_marginalization(distance_marginalization) def evaluate(self, params: dict[str, Float]) -> FloatScalar: + params = self._prepare_parameters(params) + return self._likelihood(params) + + def _prepare_parameters(self, params: dict[str, Float]) -> dict[str, Float]: + """Return the effective likelihood parameters after internal overrides.""" params = params.copy() params["trigger_time"] = self.trigger_time params["gmst"] = self.gmst @@ -280,10 +299,41 @@ def evaluate(self, params: dict[str, Float]) -> FloatScalar: if self.distance_marginalization: params["d_L"] = self.ref_dist apply_fixed_parameters(params, self.fixed_parameters) - return self._likelihood(params) + return params + + def generate_waveform( + self, params: dict[str, Float] + ) -> dict[str, Complex[Array, " n_freq"]]: + """Generate reusable sky-frame waveform polarizations. + + Distance amplitude is factored out, so the returned PyTree is a valid + cache when ``d_L`` changes. Other effective waveform inputs must remain + unchanged. + """ + prepared = self._prepare_parameters(params) + return self._generate_cached_polarizations(self.frequencies, prepared) + + def evaluate_from_waveform( + self, + params: dict[str, Float], + waveform_sky: dict[str, Complex[Array, " n_freq"]], + ) -> FloatScalar: + """Evaluate detector projection and marginalisations from a waveform cache.""" + prepared = self._prepare_parameters(params) + return self._likelihood_from_waveform(prepared, waveform_sky) def _likelihood(self, params: dict[str, Float]) -> FloatScalar: - waveform_sky = self.waveform(self.frequencies, params) + waveform_sky = self._generate_cached_polarizations(self.frequencies, params) + return self._likelihood_from_waveform(params, waveform_sky) + + def _likelihood_from_waveform( + self, + params: dict[str, Float], + waveform_sky: dict[str, Complex[Array, " n_freq"]], + ) -> FloatScalar: + """Core likelihood reduction for pre-generated waveform polarizations.""" + + waveform_sky = self._restore_cached_distance(waveform_sky, params) # --- choose accumulation type based on flags --- if self.time_marginalization: @@ -715,21 +765,59 @@ def __init__( ) def evaluate(self, params: dict[str, Float]) -> FloatScalar: + params = self._prepare_parameters(params) + return self._likelihood(params) + + def _prepare_parameters(self, params: dict[str, Float]) -> dict[str, Float]: params = params.copy() params["trigger_time"] = self.trigger_time params["gmst"] = self.gmst if self.phase_marginalization: params["phase_c"] = 0.0 apply_fixed_parameters(params, self.fixed_parameters) - return self._likelihood(params) + return params + + def generate_waveform( + self, params: dict[str, Float] + ) -> dict[str, dict[str, Complex[Array, " n_bins"]]]: + """Generate distance-normalized bin-edge polarizations for cache reuse.""" + prepared = self._prepare_parameters(params) + return { + "low": self._generate_cached_polarizations(self.freq_grid_low, prepared), + "high": self._generate_cached_polarizations( + self.freq_grid_high, prepared + ), + } + + def evaluate_from_waveform( + self, + params: dict[str, Float], + waveform_cache: dict[str, dict[str, Complex[Array, " n_bins"]]], + ) -> FloatScalar: + """Evaluate the heterodyned likelihood from cached bin-edge waveforms.""" + prepared = self._prepare_parameters(params) + return self._likelihood_from_waveform(prepared, waveform_cache) def _likelihood(self, params: dict[str, Float]) -> FloatScalar: + waveform_cache = { + "low": self._generate_cached_polarizations(self.freq_grid_low, params), + "high": self._generate_cached_polarizations(self.freq_grid_high, params), + } + return self._likelihood_from_waveform(params, waveform_cache) + + def _likelihood_from_waveform( + self, + params: dict[str, Float], + waveform_cache: dict[str, dict[str, Complex[Array, " n_bins"]]], + ) -> FloatScalar: frequencies_low = self.freq_grid_low frequencies_high = self.freq_grid_high log_likelihood: FloatScalar = jnp.zeros(()) - waveform_sky_low = self.waveform(frequencies_low, params) - waveform_sky_high = self.waveform(frequencies_high, params) + waveform_sky_low = self._restore_cached_distance(waveform_cache["low"], params) + waveform_sky_high = self._restore_cached_distance( + waveform_cache["high"], params + ) complex_d_inner_h: ComplexScalar = jnp.zeros((), dtype=jnp.complex128) diff --git a/src/jimgw/samplers/__init__.py b/src/jimgw/samplers/__init__.py index 8931416f..d49d0996 100644 --- a/src/jimgw/samplers/__init__.py +++ b/src/jimgw/samplers/__init__.py @@ -16,6 +16,7 @@ from jimgw.samplers.base import Sampler from jimgw.samplers.config import ( BaseSamplerConfig, + BlackJAXSwiGConfig, BlackJAXNSAWConfig, BlackJAXNSSConfig, BlackJAXSMCConfig, @@ -28,6 +29,7 @@ "SamplerConfig", "BaseSamplerConfig", "FlowMCConfig", + "BlackJAXSwiGConfig", "BlackJAXNSAWConfig", "BlackJAXNSSConfig", "BlackJAXSMCConfig", @@ -63,6 +65,10 @@ def build_sampler( log_likelihood_fn: Callable, log_posterior_fn: Callable, periodic: Optional[list[int] | dict[int, tuple[float, float]]] = None, + block_indices: Optional[tuple[tuple[int, ...], ...]] = None, + refresh_cache: Optional[tuple[bool, ...]] = None, + build_cache_fn: Optional[Callable] = None, + log_likelihood_from_cache_fn: Optional[Callable] = None, ) -> Sampler: """Instantiate the concrete [`Sampler`][jimgw.samplers.base.Sampler] identified by ``config.type``. @@ -86,7 +92,7 @@ def build_sampler( f"Registered types: {sorted(_REGISTRY)}" ) builder = _REGISTRY[type_str]() - return builder( + kwargs = dict( n_dims=n_dims, log_prior_fn=log_prior_fn, log_likelihood_fn=log_likelihood_fn, @@ -94,6 +100,14 @@ def build_sampler( config=config, periodic=periodic, ) + if config.type == "blackjax-swig": + kwargs.update( + block_indices=block_indices, + refresh_cache=refresh_cache, + build_cache_fn=build_cache_fn, + log_likelihood_from_cache_fn=log_likelihood_from_cache_fn, + ) + return builder(**kwargs) from jimgw.samplers.flowmc import FlowMCSampler # noqa: E402 @@ -111,3 +125,9 @@ def build_sampler( from jimgw.samplers.blackjax.nss import BlackJAXNSSSampler # noqa: E402 register_sampler("blackjax-nss", lambda: BlackJAXNSSSampler) + +from jimgw.samplers.blackjax.swig import ( # noqa: E402 + BlackJAXSwiGSampler, +) + +register_sampler("blackjax-swig", lambda: BlackJAXSwiGSampler) diff --git a/src/jimgw/samplers/blackjax/nss.py b/src/jimgw/samplers/blackjax/nss.py index c6d65e9d..d77e6363 100644 --- a/src/jimgw/samplers/blackjax/nss.py +++ b/src/jimgw/samplers/blackjax/nss.py @@ -76,6 +76,25 @@ def __init__( periodic, n_dims, sample_direction_from_covariance ) + @property + def _checkpoint_tag(self) -> str: + return "NSS" + + @property + def _update_inner_kernel_params_fn(self) -> Callable: + return live_covariance + + def _build_nested_sampler(self, n_delete: int): + config = self._config + num_inner_steps = config.num_inner_steps_per_dim * self.n_dims + return blackjax.nss( + logprior_fn=self._log_prior_fn, + loglikelihood_fn=self._log_likelihood_fn, + num_delete=n_delete, + num_inner_steps=num_inner_steps, + proposal=self._proposal, + ) + def _sample( self, rng_key: Key, @@ -101,7 +120,6 @@ def _sample( config = self._config n_live = config.n_live n_delete = int(n_live * config.n_delete_frac) - num_inner_steps = config.num_inner_steps_per_dim * self.n_dims ckpt_path = ( config.checkpoint_dir / "checkpoint.pkl" if config.checkpoint_dir is not None @@ -119,13 +137,7 @@ def _validated_initial_particles(pos): ) return arr - nested_sampler = blackjax.nss( - logprior_fn=self._log_prior_fn, - loglikelihood_fn=self._log_likelihood_fn, - num_delete=n_delete, - num_inner_steps=num_inner_steps, - proposal=self._proposal, - ) + nested_sampler = self._build_nested_sampler(n_delete) # Bypass BlackJAX's jax.vmap(init_state_fn) to avoid peak-memory OOM. # A full vmap over all live particles materialises O(n_live) concurrent @@ -145,7 +157,7 @@ def _batched_fn(pos): return _ns_adaptive_init( positions, init_state_fn=_batched_fn, - update_inner_kernel_params_fn=live_covariance, + update_inner_kernel_params_fn=self._update_inner_kernel_params_fn, ) # Resume from checkpoint if one exists. @@ -163,7 +175,10 @@ def _batched_fn(pos): n_iter = _ckpt["n_iter"] self._prev_elapsed = float(_ckpt["elapsed_time"]) logger.info( - "NSS: resumed from checkpoint at n_iter=%d (%s)", n_iter, ckpt_path + "%s: resumed from checkpoint at n_iter=%d (%s)", + self._checkpoint_tag, + n_iter, + ckpt_path, ) except ( OSError, @@ -173,7 +188,8 @@ def _batched_fn(pos): pickle.UnpicklingError, ) as _e: logger.warning( - "NSS: corrupt checkpoint at %s (%s) — starting fresh.", + "%s: corrupt checkpoint at %s (%s) — starting fresh.", + self._checkpoint_tag, ckpt_path, _e, ) @@ -214,7 +230,7 @@ def _terminate(state: AdaptiveNSState) -> bool: "elapsed_time": self._prev_elapsed + (time.perf_counter() - _method_t0), }, - "NSS", + self._checkpoint_tag, ) self._final_state = finalise(state, dead) # type: ignore[arg-type] # AdaptiveNSState structurally satisfies NSState (.particles field) diff --git a/src/jimgw/samplers/blackjax/swig.py b/src/jimgw/samplers/blackjax/swig.py new file mode 100644 index 00000000..8ae4fbaa --- /dev/null +++ b/src/jimgw/samplers/blackjax/swig.py @@ -0,0 +1,239 @@ +"""Cache-aware Nested Slice within Gibbs (SwiG) sampling.""" + +from __future__ import annotations + +from typing import Callable, NamedTuple, Optional + +import jax +import jax.numpy as jnp +from blackjax import SamplingAlgorithm +from blackjax.mcmc.slice import SliceInfo +from blackjax.mcmc.slice import build_kernel as build_slice_kernel +from blackjax.mcmc.slice import stepping_out +from blackjax.ns.from_mcmc import build_kernel as build_from_mcmc_kernel +from blackjax.ns.nss import sample_direction_from_covariance +from blackjax.smc.tuning.from_particles import particles_covariance_matrix + +from jimgw.samplers.blackjax.nss import BlackJAXNSSSampler +from jimgw.samplers.config import BlackJAXSwiGConfig +from jimgw.samplers.periodic import _build_masks_arrays + + +class CachedSliceState(NamedTuple): + """Ephemeral slice state; caches are never stored on the live particles.""" + + position: jax.Array + logdensity: jax.Array + loglikelihood: jax.Array + loglikelihood_birth: jax.Array + cache: object + + +def _build_block_covariance_update( + block_indices: tuple[tuple[int, ...], ...], +) -> Callable: + def update(rng_key, state, info, params=None): + del rng_key, info, params + covariance = jnp.atleast_2d( + particles_covariance_matrix(state.particles.position) + ) + covariances = tuple( + covariance[jnp.ix_(jnp.asarray(block), jnp.asarray(block))] + for block in block_indices + ) + return {"block_covariances": covariances} + + return update + + +def _build_swig_constrained_step( + *, + log_prior_fn: Callable, + build_cache_fn: Callable, + log_likelihood_from_cache_fn: Callable, + block_indices: tuple[tuple[int, ...], ...], + refresh_cache: tuple[bool, ...], + num_gibbs_sweeps: int, + num_inner_steps_per_dim: int, + max_steps: int, + max_shrinkage: int, + periodic: Optional[dict[int, tuple[float, float]]], + n_dims: int, +) -> Callable: + slice_kernel = build_slice_kernel( + interval=stepping_out, + max_expansions=max_steps, + max_shrinkage=max_shrinkage, + ) + periodic_mask, periodic_lower, periodic_period = _build_masks_arrays( + periodic, n_dims + ) + + def wrap(position): + return jnp.where( + periodic_mask, + periodic_lower + jnp.mod(position - periodic_lower, periodic_period), + position, + ) + + def constrained_step( + rng_key, state, loglikelihood_0, block_covariances + ): + cache = build_cache_fn(state.position) + cached_state = CachedSliceState( + position=state.position, + logdensity=state.logdensity, + loglikelihood=state.loglikelihood, + loglikelihood_birth=jnp.asarray(loglikelihood_0), + cache=cache, + ) + accepted = jnp.asarray(True) + num_expansions = jnp.asarray(0) + num_shrink = jnp.asarray(0) + + for _ in range(num_gibbs_sweeps): + for block, must_refresh, covariance in zip( + block_indices, refresh_cache, block_covariances, strict=True + ): + block_array = jnp.asarray(block) + n_steps = num_inner_steps_per_dim * len(block) + + def one_slice(carry, key): + current, all_accepted, expansions, shrink = carry + + def proposal_generator(direction_key, position, logdensity_fn): + del logdensity_fn + block_position = position[block_array] + block_direction = sample_direction_from_covariance( + direction_key, block_position, covariance + ) + direction = jnp.zeros_like(position).at[block_array].set( + block_direction + ) + + def slice_fn(t): + proposed = wrap(position + t * direction) + logprior = log_prior_fn(proposed) + if must_refresh: + proposed_cache = build_cache_fn(proposed) + else: + proposed_cache = current.cache + loglikelihood = log_likelihood_from_cache_fn( + proposed, proposed_cache + ) + proposed_state = CachedSliceState( + position=proposed, + logdensity=logprior, + loglikelihood=loglikelihood, + loglikelihood_birth=jnp.asarray(loglikelihood_0), + cache=proposed_cache, + ) + return proposed_state, loglikelihood > loglikelihood_0 + + return slice_fn + + new_state, info = slice_kernel( + key, current, None, proposal_generator + ) + return ( + new_state, + all_accepted & info.is_accepted, + expansions + info.num_expansions, + shrink + info.num_shrink, + ), None + + keys = jax.random.split(rng_key, n_steps + 1) + rng_key = keys[0] + (cached_state, accepted, num_expansions, num_shrink), _ = ( + jax.lax.scan( + one_slice, + (cached_state, accepted, num_expansions, num_shrink), + keys[1:], + ) + ) + + final_state = state._replace( + position=cached_state.position, + logdensity=cached_state.logdensity, + loglikelihood=cached_state.loglikelihood, + loglikelihood_birth=jnp.asarray(loglikelihood_0), + ) + info = SliceInfo( + is_accepted=accepted, + num_expansions=num_expansions, + num_shrink=num_shrink, + bracket_left=jnp.zeros(n_dims), + bracket_right=jnp.zeros(n_dims), + ) + return final_state, info + + return constrained_step + + +class BlackJAXSwiGSampler(BlackJAXNSSSampler): + """Nested Slice within Gibbs using cache-aware slices over named blocks.""" + + _config: BlackJAXSwiGConfig + + def __init__( + self, + *, + n_dims: int, + log_prior_fn: Callable, + log_likelihood_fn: Callable, + log_posterior_fn: Callable, + config: BlackJAXSwiGConfig, + periodic: Optional[dict[int, tuple[float, float]]] = None, + block_indices: Optional[tuple[tuple[int, ...], ...]] = None, + refresh_cache: Optional[tuple[bool, ...]] = None, + build_cache_fn: Optional[Callable] = None, + log_likelihood_from_cache_fn: Optional[Callable] = None, + ) -> None: + if block_indices is None or refresh_cache is None: + raise ValueError("resolved block indices and cache flags are required") + if build_cache_fn is None or log_likelihood_from_cache_fn is None: + raise ValueError("cache-aware likelihood callables are required") + super().__init__( + n_dims=n_dims, + log_prior_fn=log_prior_fn, + log_likelihood_fn=log_likelihood_fn, + log_posterior_fn=log_posterior_fn, + config=config, # type: ignore[arg-type] + periodic=periodic, + ) + self._block_indices = block_indices + self._refresh_cache = refresh_cache + self._build_cache_fn = build_cache_fn + self._log_likelihood_from_cache_fn = log_likelihood_from_cache_fn + self._periodic = periodic + self._block_covariance_update = _build_block_covariance_update(block_indices) + + @property + def _checkpoint_tag(self) -> str: + return "SwiG" + + @property + def _update_inner_kernel_params_fn(self) -> Callable: + return self._block_covariance_update + + def _build_nested_sampler(self, n_delete: int) -> SamplingAlgorithm: + constrained_step = _build_swig_constrained_step( + log_prior_fn=self._log_prior_fn, + build_cache_fn=self._build_cache_fn, + log_likelihood_from_cache_fn=self._log_likelihood_from_cache_fn, + block_indices=self._block_indices, + refresh_cache=self._refresh_cache, + num_gibbs_sweeps=self._config.num_gibbs_sweeps, + num_inner_steps_per_dim=self._config.num_inner_steps_per_dim, + max_steps=self._config.max_steps, + max_shrinkage=self._config.max_shrinkage, + periodic=self._periodic, + n_dims=self.n_dims, + ) + kernel = build_from_mcmc_kernel( + constrained_step, + num_inner_steps=1, + update_inner_kernel_params_fn=self._block_covariance_update, + num_delete=n_delete, + ) + return SamplingAlgorithm(lambda position, rng_key=None: position, kernel) diff --git a/src/jimgw/samplers/config.py b/src/jimgw/samplers/config.py index 2b1cb615..b68c5458 100644 --- a/src/jimgw/samplers/config.py +++ b/src/jimgw/samplers/config.py @@ -6,6 +6,7 @@ """ import logging +import math import pickle import time import warnings @@ -351,6 +352,71 @@ def _n_live_n_delete_consistency(self) -> Self: return self +class BlackJAXSwiGConfig(BaseSamplerConfig, _CheckpointMixin): + """Configuration for Nested Slice within Gibbs (SwiG) sampling. + + ``blocks`` are expressed in Jim's sampling-space parameter names. Jim + validates that they form an exact partition and determines which blocks + invalidate the waveform cache after accounting for transforms, fixed + parameters, and analytic marginalisation. + """ + + type: Literal["blackjax-swig"] = "blackjax-swig" + + blocks: list[list[str]] + n_live: int = 512 + n_delete_frac: float = 0.125 + num_gibbs_sweeps: int = 2 + num_inner_steps_per_dim: int = 1 + max_steps: int = 10 + max_shrinkage: int = 100 + termination_dlogz: float = math.exp(-3.0) + + @field_validator("blocks") + @classmethod + def _validate_blocks(cls, blocks: list[list[str]]) -> list[list[str]]: + if not blocks: + raise ValueError("blocks must contain at least one parameter block") + if any(not block for block in blocks): + raise ValueError("blocks cannot contain empty parameter blocks") + flat = [name for block in blocks for name in block] + duplicates = sorted({name for name in flat if flat.count(name) > 1}) + if duplicates: + raise ValueError(f"parameters appear in multiple blocks: {duplicates}") + return blocks + + @field_validator( + "num_gibbs_sweeps", + "num_inner_steps_per_dim", + "max_steps", + "max_shrinkage", + ) + @classmethod + def _positive_integer(cls, value: int) -> int: + if value < 1: + raise ValueError("must be >= 1") + return value + + @field_validator("n_delete_frac") + @classmethod + def _n_delete_frac_range(cls, value: float) -> float: + if not (0.0 < value < 1.0): + raise ValueError("n_delete_frac must be strictly between 0 and 1") + return value + + @model_validator(mode="after") + def _n_live_n_delete_consistency(self) -> Self: + if self.n_live < 2: + raise ValueError(f"n_live must be >= 2 (got {self.n_live}).") + n_delete = int(self.n_live * self.n_delete_frac) + if n_delete < 1: + raise ValueError( + f"n_live * n_delete_frac = {self.n_live * self.n_delete_frac} " + f"yields n_delete = {n_delete}; require n_delete >= 1." + ) + return self + + class BlackJAXSMCConfig(BaseSamplerConfig, _CheckpointMixin): """Configuration for the BlackJAX SMC sampler. @@ -460,7 +526,13 @@ def _resolve_target_ess_fraction(self) -> float: SamplerConfig = Annotated[ - Union[FlowMCConfig, BlackJAXNSAWConfig, BlackJAXNSSConfig, BlackJAXSMCConfig], + Union[ + FlowMCConfig, + BlackJAXNSAWConfig, + BlackJAXNSSConfig, + BlackJAXSwiGConfig, + BlackJAXSMCConfig, + ], Discriminator("type"), ] """Discriminated union of every concrete sampler config.""" diff --git a/tests/unit/core/single_event/test_likelihood.py b/tests/unit/core/single_event/test_likelihood.py index 01457a3a..ee9a73fc 100644 --- a/tests/unit/core/single_event/test_likelihood.py +++ b/tests/unit/core/single_event/test_likelihood.py @@ -19,7 +19,9 @@ greenwich_mean_sidereal_time as compute_gmst, ) from jimgw.core.constants import EARTH_RADIUS_LIGHT_S +from jimgw.core.jim import Jim from jimgw.core.prior import CombinePrior, GaussianPrior, PowerLawPrior, UniformPrior +from jimgw.samplers.config import BlackJAXSwiGConfig from tests.utils import assert_all_finite, common_keys_allclose FIXTURES_DIR = Path(__file__).parent.parent.parent.parent / "fixtures" @@ -101,6 +103,26 @@ def params_without_d_L_phase() -> dict: "psi": 0.0, } + @staticmethod + def swig_prior() -> CombinePrior: + bounds = { + "M_c": (20.0, 40.0), + "q": (0.5, 1.0), + "s1_z": (-0.05, 0.05), + "s2_z": (-0.05, 0.05), + "iota": (0.1, 3.0), + "ra": (0.0, 2.0 * jnp.pi), + "dec": (-1.5, 1.5), + "psi": (0.0, jnp.pi), + "t_c": (-0.05, 0.05), + } + return CombinePrior( + [ + UniformPrior(lo, hi, parameter_names=[name]) + for name, (lo, hi) in bounds.items() + ] + ) + # ── Initialization ──────────────────────────────────────────────────────── def test_initialization(self, detectors_and_waveform): @@ -114,6 +136,144 @@ def test_initialization(self, detectors_and_waveform): assert likelihood.trigger_time == gps assert hasattr(likelihood, "gmst") + def test_cached_waveform_matches_full_evaluation(self, detectors_and_waveform): + ifos, waveform, fmin, fmax, gps = detectors_and_waveform + likelihood = TransientLikelihoodFD( + detectors=ifos, + waveform=waveform, + f_min=fmin, + f_max=fmax, + trigger_time=gps, + ) + params = example_params() + cache = likelihood.generate_waveform(params) + assert jnp.allclose( + likelihood.evaluate(params), + likelihood.evaluate_from_waveform(params, cache), + ) + + def test_cached_waveform_supports_cache_reusing_changes( + self, detectors_and_waveform + ): + ifos, waveform, fmin, fmax, gps = detectors_and_waveform + likelihood = TransientLikelihoodFD( + detectors=ifos, + waveform=waveform, + f_min=fmin, + f_max=fmax, + trigger_time=gps, + phase_marginalization=True, + ) + params = example_params() + cache = likelihood.generate_waveform(params) + moved = { + **params, + "d_L": params["d_L"] * 1.2, + "ra": params["ra"] + 0.1, + "psi": 0.2, + "t_c": 0.01, + } + assert jnp.allclose( + likelihood.evaluate(moved), + likelihood.evaluate_from_waveform(moved, cache), + ) + + def test_swig_infers_waveform_blocks_after_transform_and_marginalization( + self, detectors_and_waveform + ): + ifos, waveform, fmin, fmax, gps = detectors_and_waveform + likelihood = TransientLikelihoodFD( + detectors=ifos, + waveform=waveform, + f_min=fmin, + f_max=fmax, + trigger_time=gps, + phase_marginalization=True, + distance_marginalization={"distance_prior": self.make_d_L_prior()}, + ) + blocks = [ + ["M_c", "q"], + ["s1_z", "s2_z"], + ["iota"], + ["ra", "dec"], + ["psi"], + ["t_c"], + ] + jim = Jim( + likelihood, + self.swig_prior(), + BlackJAXSwiGConfig(blocks=blocks, n_live=8, n_delete_frac=0.25), + likelihood_transforms=[MassRatioToSymmetricMassRatioTransform], + ) + assert jim.sampler._refresh_cache == (True, True, True, False, False, False) + + def test_swig_distance_block_reuses_cache_with_time_marginalization( + self, detectors_and_waveform + ): + ifos, waveform, fmin, fmax, gps = detectors_and_waveform + likelihood = TransientLikelihoodFD( + detectors=ifos, + waveform=waveform, + f_min=fmin, + f_max=fmax, + trigger_time=gps, + phase_marginalization=True, + time_marginalization={"tc_range": (-0.03, 0.03)}, + ) + prior = CombinePrior( + [ + base_prior + for base_prior in self.swig_prior().base_prior + if base_prior.parameter_names != ("t_c",) + ] + + [UniformPrior(100.0, 1000.0, parameter_names=["d_L"])] + ) + blocks = [ + ["M_c", "q"], + ["s1_z", "s2_z"], + ["iota"], + ["d_L"], + ["ra", "dec"], + ["psi"], + ] + jim = Jim( + likelihood, + prior, + BlackJAXSwiGConfig(blocks=blocks, n_live=8, n_delete_frac=0.25), + likelihood_transforms=[MassRatioToSymmetricMassRatioTransform], + ) + assert jim.sampler._refresh_cache == (True, True, True, False, False, False) + + def test_swig_rejects_marginalized_parameter_in_blocks( + self, detectors_and_waveform + ): + ifos, waveform, fmin, fmax, gps = detectors_and_waveform + likelihood = TransientLikelihoodFD( + detectors=ifos, + waveform=waveform, + f_min=fmin, + f_max=fmax, + trigger_time=gps, + phase_marginalization=True, + distance_marginalization={"distance_prior": self.make_d_L_prior()}, + ) + blocks = [ + ["M_c", "q"], + ["s1_z", "s2_z"], + ["iota"], + ["ra", "dec"], + ["psi"], + ["t_c"], + ["d_L"], + ] + with pytest.raises(ValueError, match="analytically marginalized.*d_L"): + Jim( + likelihood, + self.swig_prior(), + BlackJAXSwiGConfig(blocks=blocks, n_live=8, n_delete_frac=0.25), + likelihood_transforms=[MassRatioToSymmetricMassRatioTransform], + ) + def test_uninitialized_data_raises(self): gps = 1126259462.4 ifos = [get_H1(), get_L1()] @@ -1071,6 +1231,50 @@ def test_evaluate_jit_matches(self, detectors_and_waveform): likelihood.evaluate(params), jax.jit(likelihood.evaluate)(params) ) + def test_cached_waveform_matches_full_evaluation(self, detectors_and_waveform): + ifos, waveform, fmin, fmax, gps = detectors_and_waveform + likelihood = HeterodynedTransientLikelihoodFD( + detectors=ifos, + waveform=waveform, + f_min=fmin, + f_max=fmax, + trigger_time=gps, + n_bins=32, + reference_parameters=example_params(), + ) + params = example_params() + cache = likelihood.generate_waveform(params) + assert jnp.allclose( + likelihood.evaluate_from_waveform(params, cache), + likelihood.evaluate(params), + ) + + def test_cached_waveform_supports_cache_reusing_changes( + self, detectors_and_waveform + ): + ifos, waveform, fmin, fmax, gps = detectors_and_waveform + likelihood = HeterodynedTransientLikelihoodFD( + detectors=ifos, + waveform=waveform, + f_min=fmin, + f_max=fmax, + trigger_time=gps, + n_bins=32, + reference_parameters=example_params(), + ) + params = example_params() + cache = likelihood.generate_waveform(params) + projected = { + **params, + "d_L": params["d_L"] * 1.2, + "psi": params["psi"] + 0.1, + "t_c": 0.002, + } + assert jnp.allclose( + likelihood.evaluate_from_waveform(projected, cache), + likelihood.evaluate(projected), + ) + def test_evaluate_different_fmin(self, detectors_and_waveform): ifos, waveform, fmin, fmax, gps = detectors_and_waveform likelihood = HeterodynedTransientLikelihoodFD( diff --git a/tests/unit/samplers/blackjax/test_swig.py b/tests/unit/samplers/blackjax/test_swig.py new file mode 100644 index 00000000..df5232f0 --- /dev/null +++ b/tests/unit/samplers/blackjax/test_swig.py @@ -0,0 +1,62 @@ +"""Tests for cache-aware Nested Slice within Gibbs.""" + +import jax +import jax.numpy as jnp +import numpy as np + +from jimgw.samplers.blackjax.swig import BlackJAXSwiGSampler +from jimgw.samplers.config import BlackJAXSwiGConfig + + +def _log_prior(position): + return jnp.where(jnp.all((position >= 0.0) & (position <= 1.0)), 0.0, -jnp.inf) + + +def _build_cache(position): + return position[0] ** 2 + + +def _log_likelihood_from_cache(position, cache): + return -40.0 * ((cache - 0.25) ** 2 + (position[1] - 0.5) ** 2) + + +def _log_likelihood(position): + return _log_likelihood_from_cache(position, _build_cache(position)) + + +def _make_sampler() -> BlackJAXSwiGSampler: + config = BlackJAXSwiGConfig( + blocks=[["slow"], ["fast"]], + n_live=24, + n_delete_frac=0.25, + termination_dlogz=1.5, + max_steps=4, + max_shrinkage=30, + ) + return BlackJAXSwiGSampler( + n_dims=2, + log_prior_fn=_log_prior, + log_likelihood_fn=_log_likelihood, + log_posterior_fn=lambda x: _log_prior(x) + _log_likelihood(x), + config=config, + block_indices=((0,), (1,)), + refresh_cache=(True, False), + build_cache_fn=_build_cache, + log_likelihood_from_cache_fn=_log_likelihood_from_cache, + ) + + +def test_swig_cached_likelihood_remains_consistent(): + sampler = _make_sampler() + initial = jax.random.uniform(jax.random.key(1), (24, 2)) + sampler.sample(jax.random.key(2), initial) + result = sampler.get_samples() + expected = jax.vmap(_log_likelihood)(jnp.asarray(result["samples"])) + np.testing.assert_allclose(result["log_likelihood"], expected, rtol=1e-10) + + +def test_swig_does_not_store_cache_on_live_particles(): + sampler = _make_sampler() + initial = jax.random.uniform(jax.random.key(3), (24, 2)) + sampler.sample(jax.random.key(4), initial) + assert not hasattr(sampler._final_state.particles, "cache") diff --git a/tests/unit/samplers/test_config.py b/tests/unit/samplers/test_config.py index d8cde16f..38d4918a 100644 --- a/tests/unit/samplers/test_config.py +++ b/tests/unit/samplers/test_config.py @@ -1,4 +1,5 @@ import warnings +import math import jax import numpy as np @@ -9,6 +10,7 @@ BlackJAXNSAWConfig, BlackJAXNSSConfig, BlackJAXSMCConfig, + BlackJAXSwiGConfig, FlowMCConfig, GRWConfig, HMCConfig, @@ -46,6 +48,26 @@ def test_sampler_config_union_from_dict(): cfg4 = ta.validate_python({"type": "blackjax-smc"}) assert isinstance(cfg4, BlackJAXSMCConfig) + cfg5 = ta.validate_python( + {"type": "blackjax-swig", "blocks": [["x"], ["y"]]} + ) + assert isinstance(cfg5, BlackJAXSwiGConfig) + + +def test_swig_blocks_must_be_nonempty_and_unique(): + with pytest.raises(ValidationError, match="at least one"): + BlackJAXSwiGConfig(blocks=[]) + with pytest.raises(ValidationError, match="empty"): + BlackJAXSwiGConfig(blocks=[["x"], []]) + with pytest.raises(ValidationError, match="multiple blocks"): + BlackJAXSwiGConfig(blocks=[["x"], ["x"]]) + + +def test_swig_sampling_defaults(): + config = BlackJAXSwiGConfig(blocks=[["x"]]) + assert config.num_gibbs_sweeps == 2 + assert config.termination_dlogz == pytest.approx(math.exp(-3.0)) + def test_extra_fields_forbidden(): with pytest.raises(ValidationError): diff --git a/tests/unit/samplers/test_registry.py b/tests/unit/samplers/test_registry.py index a03c475e..d92ec9aa 100644 --- a/tests/unit/samplers/test_registry.py +++ b/tests/unit/samplers/test_registry.py @@ -77,10 +77,11 @@ class _FakeConfig(BaseSamplerConfig): ) -def test_registry_has_all_four_types(): +def test_registry_has_all_sampler_types(): from jimgw.samplers import _REGISTRY assert "flowmc" in _REGISTRY assert "blackjax-ns-aw" in _REGISTRY assert "blackjax-nss" in _REGISTRY + assert "blackjax-swig" in _REGISTRY assert "blackjax-smc" in _REGISTRY From d1b5cd9eae3bf2a4fce2dffadb0f057637107c1c Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Thu, 16 Jul 2026 11:33:01 +0000 Subject: [PATCH 2/2] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- src/jimgw/core/single_event/likelihood.py | 9 +++++---- src/jimgw/samplers/blackjax/swig.py | 20 +++++++++----------- 2 files changed, 14 insertions(+), 15 deletions(-) diff --git a/src/jimgw/core/single_event/likelihood.py b/src/jimgw/core/single_event/likelihood.py index 8932689c..ba5f5ffa 100644 --- a/src/jimgw/core/single_event/likelihood.py +++ b/src/jimgw/core/single_event/likelihood.py @@ -104,7 +104,10 @@ def _restore_cached_distance(self, waveform_sky, params): if "d_L" not in getattr(self.waveform, "parameter_names", ()): return waveform_sky scale = 1.0 / params["d_L"] - return {polarization: strain * scale for polarization, strain in waveform_sky.items()} + return { + polarization: strain * scale + for polarization, strain in waveform_sky.items() + } def evaluate(self, params: dict[str, Float]) -> FloatScalar: """Apply ``fixed_parameters`` overrides and evaluate the likelihood. @@ -784,9 +787,7 @@ def generate_waveform( prepared = self._prepare_parameters(params) return { "low": self._generate_cached_polarizations(self.freq_grid_low, prepared), - "high": self._generate_cached_polarizations( - self.freq_grid_high, prepared - ), + "high": self._generate_cached_polarizations(self.freq_grid_high, prepared), } def evaluate_from_waveform( diff --git a/src/jimgw/samplers/blackjax/swig.py b/src/jimgw/samplers/blackjax/swig.py index 8ae4fbaa..e1424c8a 100644 --- a/src/jimgw/samplers/blackjax/swig.py +++ b/src/jimgw/samplers/blackjax/swig.py @@ -76,9 +76,7 @@ def wrap(position): position, ) - def constrained_step( - rng_key, state, loglikelihood_0, block_covariances - ): + def constrained_step(rng_key, state, loglikelihood_0, block_covariances): cache = build_cache_fn(state.position) cached_state = CachedSliceState( position=state.position, @@ -107,8 +105,10 @@ def proposal_generator(direction_key, position, logdensity_fn): block_direction = sample_direction_from_covariance( direction_key, block_position, covariance ) - direction = jnp.zeros_like(position).at[block_array].set( - block_direction + direction = ( + jnp.zeros_like(position) + .at[block_array] + .set(block_direction) ) def slice_fn(t): @@ -144,12 +144,10 @@ def slice_fn(t): keys = jax.random.split(rng_key, n_steps + 1) rng_key = keys[0] - (cached_state, accepted, num_expansions, num_shrink), _ = ( - jax.lax.scan( - one_slice, - (cached_state, accepted, num_expansions, num_shrink), - keys[1:], - ) + (cached_state, accepted, num_expansions, num_shrink), _ = jax.lax.scan( + one_slice, + (cached_state, accepted, num_expansions, num_shrink), + keys[1:], ) final_state = state._replace(