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..8b5b9ddb 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,9 +106,13 @@ 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 an equal number of samples per posterior. + Set to :code:`True` for backward compatibility. """ self.samples_per_posterior = max_samples + self.equal_samples = require_equal_samples self.data = self.resample_posteriors(posteriors, max_samples=max_samples) if isinstance(hyper_prior, types.FunctionType): @@ -195,9 +200,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**2) variance = (square_expectation - expectation**2) / ( self.samples_per_posterior * expectation**2 ) @@ -205,6 +210,19 @@ 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)]) + transitions = xp.concat( + [xp.zeros(1), xp.cumsum(self.samples_per_posterior)] + ).astype(int) + expectation = ( + cumulative[transitions[1:]] - cumulative[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 @@ -320,8 +338,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 ------- @@ -329,17 +346,33 @@ 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 + 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) + for key in data: - data[key] = xp.array(data[key]) + data[key] = xp.asarray(data[key]) + return data def posterior_predictive_resample(self, samples, return_weights=False): @@ -370,47 +403,83 @@ 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: + 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() 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: 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/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..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"]) @@ -86,10 +88,22 @@ 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)) + 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) @@ -204,6 +218,21 @@ 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: + 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: @@ -216,6 +245,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"],