Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions gwpopulation/experimental/numpyro.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
131 changes: 100 additions & 31 deletions gwpopulation/hyperpe.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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.
Comment thread
ColmTalbot marked this conversation as resolved.
"""

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):
Expand Down Expand Up @@ -195,16 +200,29 @@ 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
)
return xp.log(expectation), variance
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
Expand Down Expand Up @@ -320,26 +338,41 @@ 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
-------
data: dict
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
Comment thread
ColmTalbot marked this conversation as resolved.
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):
Expand Down Expand Up @@ -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}")
Expand Down
7 changes: 5 additions & 2 deletions test/example_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")

Expand Down Expand Up @@ -67,6 +69,7 @@ def test_likelihood_evaluation(backend, jit):
hyper_prior=model,
posteriors=posteriors,
selection_function=selection,
require_equal_samples=equal,
)
Comment thread
ColmTalbot marked this conversation as resolved.

priors = bilby.core.prior.PriorDict("priors/bbh_population.prior")
Expand Down
32 changes: 31 additions & 1 deletion test/likelihood_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"])
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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:
Expand All @@ -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=["<lambda>", "SinglePeakSmoothedMassDistribution"],
Expand Down
Loading