From 4ca7d6a82eedccda749dcb7442a2fded38e1ce75 Mon Sep 17 00:00:00 2001 From: Colm Talbot Date: Thu, 16 Jul 2026 21:12:57 +0000 Subject: [PATCH 1/4] FEAT: allow use of different numbers of samples per event --- gwpopulation/experimental/numpyro.py | 4 +- gwpopulation/hyperpe.py | 129 ++++++++++++++++++++------- test/example_test.py | 7 +- test/likelihood_test.py | 4 +- 4 files changed, 108 insertions(+), 36 deletions(-) diff --git a/gwpopulation/experimental/numpyro.py b/gwpopulation/experimental/numpyro.py index 9c55d93c..4249bbd3 100644 --- a/gwpopulation/experimental/numpyro.py +++ b/gwpopulation/experimental/numpyro.py @@ -43,8 +43,8 @@ def gwpopulation_likelihood_model( likelihood.hyper_prior.prob(likelihood.data, **parameters) / likelihood.sampling_prior ) - expectations = jnp.mean(weights, axis=-1) - square_expectations = jnp.mean(weights**2, axis=-1) + expectations = likelihood._weight_expectation(weights) + square_expectations = likelihood._weight_expectation(weights**2) variances = deterministic( "variances", (square_expectations - expectations**2) diff --git a/gwpopulation/hyperpe.py b/gwpopulation/hyperpe.py index 7451ef83..8351e6f7 100644 --- a/gwpopulation/hyperpe.py +++ b/gwpopulation/hyperpe.py @@ -78,6 +78,7 @@ def __init__( selection_function=lambda args: 1, conversion_function=lambda args: (args, None), maximum_uncertainty=xp.inf, + require_equal_samples=False, ): """ Parameters @@ -105,10 +106,16 @@ def __init__( The maximum allowed uncertainty in the natural log likelihood. If the uncertainty is larger than this value a log likelihood of -inf will be returned. Default = inf + require_equal_samples: bool + Whether to require and equal number of samples per posterior. + Set to :code:`True` for backward compatibility. """ self.samples_per_posterior = max_samples - self.data = self.resample_posteriors(posteriors, max_samples=max_samples) + self.equal_samples = require_equal_samples + self.data, self.transitions = self.resample_posteriors( + posteriors, max_samples=max_samples + ) if isinstance(hyper_prior, types.FunctionType): hyper_prior = Model([hyper_prior]) @@ -195,9 +202,9 @@ def _compute_per_event_ln_bayes_factors( self, parameters, *, return_uncertainty=True ): weights = self.hyper_prior.prob(self.data, **parameters) / self.sampling_prior - expectation = xp.mean(weights, axis=-1) + expectation = self._weight_expectation(weights) if return_uncertainty: - square_expectation = xp.mean(weights**2, axis=-1) + square_expectation = self._weight_expectation(weights) variance = (square_expectation - expectation**2) / ( self.samples_per_posterior * expectation**2 ) @@ -205,6 +212,17 @@ def _compute_per_event_ln_bayes_factors( else: return xp.log(expectation) + def _weight_expectation(self, weights): + if self.equal_samples: + expectation = xp.mean(weights, axis=-1) + else: + cumulative = xp.concat([xp.zeros(1), xp.cumsum(weights)]) + expectation = ( + cumulative[self.transitions[1:]] + - cumulative[self.transitions[:-1]] + ) / self.samples_per_posterior + return expectation + def _get_selection_factor(self, parameters, *, return_uncertainty=True): selection, variance = self._selection_function_with_uncertainty( parameters=parameters @@ -329,18 +347,36 @@ def resample_posteriors(self, posteriors, max_samples=1e300): Dictionary containing arrays of size (n_posteriors, max_samples) There is a key for each shared key in posteriors. """ - for posterior in posteriors: - max_samples = min(len(posterior), max_samples) data = {key: [] for key in posteriors[0]} - logger.debug(f"Downsampling to {max_samples} samples per posterior.") - self.samples_per_posterior = max_samples - for posterior in posteriors: - temp = posterior.sample(self.samples_per_posterior) - for key in data: - data[key].append(temp[key]) + + if self.equal_samples: + for posterior in posteriors: + max_samples = min(len(posterior), max_samples) + logger.debug(f"Downsampling to {max_samples} samples per posterior.") + self.samples_per_posterior = max_samples + transitions = None + for posterior in posteriors: + temp = posterior.sample(self.samples_per_posterior) + for key in data: + data[key].append(temp[key]) + else: + self.samples_per_posterior = np.asarray([ + min(len(posterior), max_samples) for posterior in posteriors + ]) + transitions = np.concat([ + np.zeros(1), np.cumsum(self.samples_per_posterior) + ]).astype(int) + for posterior, nsamples in zip(posteriors, self.samples_per_posterior): + temp = posterior.sample(nsamples) + for key in data: + data[key].extend(temp[key]) + self.samples_per_posterior = xp.asarray(self.samples_per_posterior) + transitions = xp.asarray(transitions) + for key in data: - data[key] = xp.array(data[key]) - return data + data[key] = xp.asarray(data[key]) + + return data, transitions def posterior_predictive_resample(self, samples, return_weights=False): """ @@ -370,47 +406,78 @@ def posterior_predictive_resample(self, samples, return_weights=False): samples = [dict(samples.iloc[ii]) for ii in range(len(samples))] elif isinstance(samples, dict): samples = [samples] - weights = xp.zeros((self.n_posteriors, self.samples_per_posterior)) + if self.equal_samples: + weights = xp.zeros((self.n_posteriors, self.samples_per_posterior)) + else: + weights = xp.zeros(int(xp.sum(self.samples_per_posterior))) + event_weights = xp.zeros(self.n_posteriors) for sample in tqdm(samples): parameters, added_keys = self.conversion_function(sample.copy()) new_weights = ( self.hyper_prior.prob(self.data, **parameters) / self.sampling_prior ) - event_weights += xp.mean(new_weights, axis=-1) - new_weights = (new_weights.T / xp.sum(new_weights, axis=-1)).T + expectation = self._weight_expectation(new_weights) + event_weights += expectation + if self.equal_samples: + denominator = expectation * self.samples_per_posterior + else: + denominator = xp.concat([ + xp.ones(nsamples) * weight * nsamples + for nsamples, weight in zip(self.samples_per_posterior, expectation) + ]) + new_weights = (new_weights.T / denominator).T weights += new_weights - weights = (weights.T / xp.sum(weights, axis=-1)).T + new_idxs = xp.empty_like(weights, dtype=int) for ii in range(self.n_posteriors): + if self.equal_samples: + sl = ii + start = 0 + nsamples = self.samples_per_posterior + else: + sl = slice(self.transitions[ii], self.transitions[ii + 1]) + start = self.transitions[ii] + nsamples = int(self.samples_per_posterior[ii]) + wts = weights[sl] + wts /= wts.sum() if "jax" in xp.__name__: from jax import random rng_key = random.PRNGKey(np.random.randint(10000000)) - new_idxs = new_idxs.at[ii].set( + new_idxs = new_idxs.at[sl].set( random.choice( rng_key, - xp.arange(self.samples_per_posterior), - shape=(self.samples_per_posterior,), + xp.arange(nsamples) + start, + shape=(nsamples,), replace=True, - p=weights[ii], + p=wts, ) ) else: - new_idxs[ii] = xp.asarray( + new_idxs[sl] = xp.asarray( np.random.choice( - range(self.samples_per_posterior), - size=self.samples_per_posterior, + np.arange(nsamples) + start, + size=nsamples, replace=True, - p=to_numpy(weights[ii]), + p=to_numpy(wts), ) ) - new_samples = { - key: xp.vstack( - [self.data[key][ii, new_idxs[ii]] for ii in range(self.n_posteriors)] - ) - for key in self.data - } + + if self.equal_samples: + new_samples = { + key: xp.vstack( + [self.data[key][ii, new_idxs[ii]] for ii in range(self.n_posteriors)] + ) + for key in self.data + } + else: + new_samples = { + key: xp.vstack( + [self.data[key][new_idxs] for ii in range(self.n_posteriors)] + ) + for key in self.data + } event_weights = list(event_weights) weight_string = " ".join([f"{float(weight):.1f}" for weight in event_weights]) logger.info(f"Resampling done, sum of weights for events are {weight_string}") diff --git a/test/example_test.py b/test/example_test.py index f769594b..74e4df87 100644 --- a/test/example_test.py +++ b/test/example_test.py @@ -9,8 +9,10 @@ from gwpopulation.experimental.jax import JittedLikelihood -@pytest.mark.parametrize("jit", [True, False]) -def test_likelihood_evaluation(backend, jit): +@pytest.mark.parametrize( + "jit, equal", [[True, True], [True, False], [False, True], [False, False]] +) +def test_likelihood_evaluation(backend, jit, equal): if backend != "jax" and jit: pytest.skip(reason="JIT only works with JAX") @@ -67,6 +69,7 @@ def test_likelihood_evaluation(backend, jit): hyper_prior=model, posteriors=posteriors, selection_function=selection, + require_equal_samples=equal, ) priors = bilby.core.prior.PriorDict("priors/bbh_population.prior") diff --git a/test/likelihood_test.py b/test/likelihood_test.py index d4a0ae09..2f4f14e1 100644 --- a/test/likelihood_test.py +++ b/test/likelihood_test.py @@ -86,7 +86,7 @@ def test_hpe_likelihood_set_selection(self): def test_hpe_likelihood_set_max_samples(self): like = HyperparameterLikelihood( - posteriors=self.data, hyper_prior=self.model, max_samples=10 + posteriors=self.data, hyper_prior=self.model, max_samples=10, require_equal_samples=True ) self.assertEqual(like.data["a"].shape, (5, 10)) @@ -204,6 +204,7 @@ def test_resampling_posteriors(self): hyper_prior=self.model, selection_function=self.selection_function, ln_evidences=self.ln_evidences, + require_equal_samples=True, ) new_samples = like.posterior_predictive_resample(samples=samples) for key in new_samples: @@ -216,6 +217,7 @@ def test_meta_data(self): hyper_prior=model, selection_function=self.selection_function, ln_evidences=self.ln_evidences, + require_equal_samples=True, ) expected = dict( model=["", "SinglePeakSmoothedMassDistribution"], From 39c72816653e877422f3c79040d284593c75527b Mon Sep 17 00:00:00 2001 From: Colm Talbot Date: Fri, 17 Jul 2026 13:32:03 +0000 Subject: [PATCH 2/4] Formatting fixes --- gwpopulation/hyperpe.py | 34 ++++++++++++++++++++-------------- test/likelihood_test.py | 5 ++++- 2 files changed, 24 insertions(+), 15 deletions(-) diff --git a/gwpopulation/hyperpe.py b/gwpopulation/hyperpe.py index 8351e6f7..339e4983 100644 --- a/gwpopulation/hyperpe.py +++ b/gwpopulation/hyperpe.py @@ -204,7 +204,7 @@ def _compute_per_event_ln_bayes_factors( weights = self.hyper_prior.prob(self.data, **parameters) / self.sampling_prior expectation = self._weight_expectation(weights) if return_uncertainty: - square_expectation = self._weight_expectation(weights) + square_expectation = self._weight_expectation(weights**2) variance = (square_expectation - expectation**2) / ( self.samples_per_posterior * expectation**2 ) @@ -218,8 +218,7 @@ def _weight_expectation(self, weights): else: cumulative = xp.concat([xp.zeros(1), xp.cumsum(weights)]) expectation = ( - cumulative[self.transitions[1:]] - - cumulative[self.transitions[:-1]] + cumulative[self.transitions[1:]] - cumulative[self.transitions[:-1]] ) / self.samples_per_posterior return expectation @@ -360,12 +359,12 @@ def resample_posteriors(self, posteriors, max_samples=1e300): for key in data: data[key].append(temp[key]) else: - self.samples_per_posterior = np.asarray([ - min(len(posterior), max_samples) for posterior in posteriors - ]) - transitions = np.concat([ - np.zeros(1), np.cumsum(self.samples_per_posterior) - ]).astype(int) + self.samples_per_posterior = np.asarray( + [min(len(posterior), max_samples) for posterior in posteriors] + ) + transitions = np.concat( + [np.zeros(1), np.cumsum(self.samples_per_posterior)] + ).astype(int) for posterior, nsamples in zip(posteriors, self.samples_per_posterior): temp = posterior.sample(nsamples) for key in data: @@ -422,10 +421,14 @@ def posterior_predictive_resample(self, samples, return_weights=False): if self.equal_samples: denominator = expectation * self.samples_per_posterior else: - denominator = xp.concat([ - xp.ones(nsamples) * weight * nsamples - for nsamples, weight in zip(self.samples_per_posterior, expectation) - ]) + denominator = xp.concat( + [ + xp.ones(nsamples) * weight * nsamples + for nsamples, weight in zip( + self.samples_per_posterior, expectation + ) + ] + ) new_weights = (new_weights.T / denominator).T weights += new_weights @@ -467,7 +470,10 @@ def posterior_predictive_resample(self, samples, return_weights=False): if self.equal_samples: new_samples = { key: xp.vstack( - [self.data[key][ii, new_idxs[ii]] for ii in range(self.n_posteriors)] + [ + self.data[key][ii, new_idxs[ii]] + for ii in range(self.n_posteriors) + ] ) for key in self.data } diff --git a/test/likelihood_test.py b/test/likelihood_test.py index 2f4f14e1..a7eef65b 100644 --- a/test/likelihood_test.py +++ b/test/likelihood_test.py @@ -86,7 +86,10 @@ def test_hpe_likelihood_set_selection(self): def test_hpe_likelihood_set_max_samples(self): like = HyperparameterLikelihood( - posteriors=self.data, hyper_prior=self.model, max_samples=10, require_equal_samples=True + posteriors=self.data, + hyper_prior=self.model, + max_samples=10, + require_equal_samples=True, ) self.assertEqual(like.data["a"].shape, (5, 10)) From 8d4eff780ddfe671c260e7fb2d6a57b0d6fb5dc7 Mon Sep 17 00:00:00 2001 From: Colm Talbot Date: Fri, 17 Jul 2026 14:02:37 +0000 Subject: [PATCH 3/4] Add testing of unequal sample number case --- gwpopulation/hyperpe.py | 12 +++--------- test/likelihood_test.py | 25 +++++++++++++++++++++++++ 2 files changed, 28 insertions(+), 9 deletions(-) diff --git a/gwpopulation/hyperpe.py b/gwpopulation/hyperpe.py index 339e4983..2b2f3dc2 100644 --- a/gwpopulation/hyperpe.py +++ b/gwpopulation/hyperpe.py @@ -107,7 +107,7 @@ def __init__( If the uncertainty is larger than this value a log likelihood of -inf will be returned. Default = inf require_equal_samples: bool - Whether to require and equal number of samples per posterior. + Whether to require an equal number of samples per posterior. Set to :code:`True` for backward compatibility. """ @@ -337,8 +337,7 @@ def resample_posteriors(self, posteriors, max_samples=1e300): posteriors: list List of pandas DataFrame objects. max_samples: int, opt - Maximum number of samples to take from each posterior, - default is length of shortest posterior chain. + Maximum number of samples to take from each posterior. Returns ------- @@ -478,12 +477,7 @@ def posterior_predictive_resample(self, samples, return_weights=False): for key in self.data } else: - new_samples = { - key: xp.vstack( - [self.data[key][new_idxs] for ii in range(self.n_posteriors)] - ) - for key in self.data - } + new_samples = {key: self.data[key][new_idxs] for key in self.data} event_weights = list(event_weights) weight_string = " ".join([f"{float(weight):.1f}" for weight in event_weights]) logger.info(f"Resampling done, sum of weights for events are {weight_string}") diff --git a/test/likelihood_test.py b/test/likelihood_test.py index a7eef65b..73469c93 100644 --- a/test/likelihood_test.py +++ b/test/likelihood_test.py @@ -23,6 +23,8 @@ def setUp(self): self.model = lambda dataset, a, b, c: dataset["a"] one_data = pd.DataFrame({key: xp.ones(500) for key in self.params}) self.data = [one_data] * 5 + self.sample_lengths = [100, 123, 459, 43, 233] + self.unequal_data = [one_data[:ns] for ns in self.sample_lengths] self.ln_evidences = [0] * 5 self.selection_function = lambda args: 2.0 self.conversion_function = lambda args: (args, ["bar"]) @@ -93,6 +95,15 @@ def test_hpe_likelihood_set_max_samples(self): ) self.assertEqual(like.data["a"].shape, (5, 10)) + def test_hpe_likelihood_unequal_samples(self): + like = HyperparameterLikelihood( + posteriors=self.unequal_data, + hyper_prior=self.model, + require_equal_samples=False, + ) + for value in like.data.values(): + self.assertEqual(value.shape, (sum(self.sample_lengths),)) + def test_hpe_likelihood_log_likelihood_ratio(self): like = HyperparameterLikelihood(posteriors=self.data, hyper_prior=self.model) self.assertEqual(like.log_likelihood_ratio(self.params), 0.0) @@ -213,6 +224,20 @@ def test_resampling_posteriors(self): for key in new_samples: self.assertEqual(new_samples[key].shape, like.data[key].shape) + def test_resampling_unequal_posteriors(self): + priors = PriorDict(dict(a=Uniform(0, 2), b=Uniform(0, 2), c=Uniform(0, 2))) + samples = priors.sample(100) + like = HyperparameterLikelihood( + posteriors=self.unequal_data, + hyper_prior=self.model, + selection_function=self.selection_function, + ln_evidences=self.ln_evidences, + require_equal_samples=False, + ) + new_samples = like.posterior_predictive_resample(samples=samples) + for key in new_samples: + self.assertEqual(new_samples[key].shape, like.data[key].shape) + def test_meta_data(self): model = Model([self.model, SinglePeakSmoothedMassDistribution()]) like = HyperparameterLikelihood( From 3b8045d1acd59e94c7bf7173960631cd82bea1fb Mon Sep 17 00:00:00 2001 From: Colm Talbot Date: Fri, 17 Jul 2026 14:27:49 +0000 Subject: [PATCH 4/4] Dont cache transition points --- gwpopulation/hyperpe.py | 20 +++++++++++--------- 1 file changed, 11 insertions(+), 9 deletions(-) diff --git a/gwpopulation/hyperpe.py b/gwpopulation/hyperpe.py index 2b2f3dc2..8b5b9ddb 100644 --- a/gwpopulation/hyperpe.py +++ b/gwpopulation/hyperpe.py @@ -113,9 +113,7 @@ def __init__( self.samples_per_posterior = max_samples self.equal_samples = require_equal_samples - self.data, self.transitions = self.resample_posteriors( - posteriors, max_samples=max_samples - ) + self.data = self.resample_posteriors(posteriors, max_samples=max_samples) if isinstance(hyper_prior, types.FunctionType): hyper_prior = Model([hyper_prior]) @@ -217,8 +215,11 @@ def _weight_expectation(self, weights): expectation = xp.mean(weights, axis=-1) else: cumulative = xp.concat([xp.zeros(1), xp.cumsum(weights)]) + transitions = xp.concat( + [xp.zeros(1), xp.cumsum(self.samples_per_posterior)] + ).astype(int) expectation = ( - cumulative[self.transitions[1:]] - cumulative[self.transitions[:-1]] + cumulative[transitions[1:]] - cumulative[transitions[:-1]] ) / self.samples_per_posterior return expectation @@ -352,7 +353,6 @@ def resample_posteriors(self, posteriors, max_samples=1e300): max_samples = min(len(posterior), max_samples) logger.debug(f"Downsampling to {max_samples} samples per posterior.") self.samples_per_posterior = max_samples - transitions = None for posterior in posteriors: temp = posterior.sample(self.samples_per_posterior) for key in data: @@ -369,12 +369,11 @@ def resample_posteriors(self, posteriors, max_samples=1e300): for key in data: data[key].extend(temp[key]) self.samples_per_posterior = xp.asarray(self.samples_per_posterior) - transitions = xp.asarray(transitions) for key in data: data[key] = xp.asarray(data[key]) - return data, transitions + return data def posterior_predictive_resample(self, samples, return_weights=False): """ @@ -438,8 +437,11 @@ def posterior_predictive_resample(self, samples, return_weights=False): start = 0 nsamples = self.samples_per_posterior else: - sl = slice(self.transitions[ii], self.transitions[ii + 1]) - start = self.transitions[ii] + transitions = np.concat( + [xp.zeros(1), xp.cumsum(self.samples_per_posterior)] + ).astype(int) + sl = slice(transitions[ii], transitions[ii + 1]) + start = transitions[ii] nsamples = int(self.samples_per_posterior[ii]) wts = weights[sl] wts /= wts.sum()