Skip to content
Open
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
42 changes: 20 additions & 22 deletions bilby/bilby_mcmc/sampler.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,6 @@
MCMCSampler,
ResumeError,
SamplerError,
_sampling_convenience_dump,
signal_wrapper,
)
from ..core.utils import (
Expand All @@ -24,6 +23,7 @@
random,
safe_file_dump,
)
from ..core.utils.parallel import sampling_convenience_dump
from . import proposals
from .chain import Chain, Sample
from .utils import LOGLKEY, LOGPKEY, ConvergenceInputs, ParallelTemperingInputs
Expand Down Expand Up @@ -284,7 +284,7 @@ def add_data_to_result(result, ptsampler, outdir, label, make_plots):
total_steps=ptsampler.position,
nsamples=ptsampler.nsamples,
)
if ptsampler.pool is not None:
if ptsampler.pool is not None and hasattr(ptsampler.pool, "_processes"):
npool = ptsampler.pool._processes
else:
npool = 1
Expand Down Expand Up @@ -615,7 +615,7 @@ def __init__(

self._nsamples_dict = {}
self.ensemble_proposal_cycle = proposals.get_default_ensemble_proposal_cycle(
_sampling_convenience_dump.priors
sampling_convenience_dump.priors
)
self.sampling_time = 0
self.ln_z_dict = dict()
Expand All @@ -630,7 +630,7 @@ def get_initial_betas(self):
elif pt_inputs.Tmax is not None:
betas = np.logspace(0, -np.log10(pt_inputs.Tmax), pt_inputs.ntemps)
elif pt_inputs.Tmax_from_SNR is not None:
ndim = len(_sampling_convenience_dump.priors.non_fixed_keys)
ndim = len(sampling_convenience_dump.priors.non_fixed_keys)
target_hot_likelihood = ndim / 2
Tmax = pt_inputs.Tmax_from_SNR**2 / (2 * target_hot_likelihood)
betas = np.logspace(0, -np.log10(Tmax), pt_inputs.ntemps)
Expand Down Expand Up @@ -1172,15 +1172,15 @@ def __init__(
self.Eindex = Eindex
self.use_ratio = use_ratio
self.normalize_prior = normalize_prior
self.parameters = _sampling_convenience_dump.priors.non_fixed_keys
self.parameters = sampling_convenience_dump.priors.non_fixed_keys
self.ndim = len(self.parameters)

if initial_sample_method.lower() == "prior":
full_sample_dict = _sampling_convenience_dump.priors.sample()
full_sample_dict = sampling_convenience_dump.priors.sample()
initial_sample = {
k: v
for k, v in full_sample_dict.items()
if k in _sampling_convenience_dump.priors.non_fixed_keys
if k in sampling_convenience_dump.priors.non_fixed_keys
}
elif initial_sample_method.lower() in ["maximize", "maximise", "maximum"]:
initial_sample = get_initial_maximimum_posterior_sample(self.beta)
Expand Down Expand Up @@ -1217,7 +1217,7 @@ def __init__(

self.proposal_cycle = proposals.get_proposal_cycle(
proposal_cycle,
_sampling_convenience_dump.priors,
sampling_convenience_dump.priors,
L1steps=self.chain.L1steps,
warn=warn,
)
Expand All @@ -1236,16 +1236,16 @@ def set_convergence_inputs(self, convergence_inputs):
self.stop_after_convergence = convergence_inputs.stop_after_convergence

def log_likelihood(self, sample):
params = deepcopy(_sampling_convenience_dump.parameters)
params = deepcopy(sampling_convenience_dump.parameters)
params.update(sample.sample_dict)

if self.use_ratio:
return _sampling_convenience_dump.likelihood.log_likelihood_ratio(params)
return sampling_convenience_dump.likelihood.log_likelihood_ratio(params)
else:
return _sampling_convenience_dump.likelihood.log_likelihood(params)
return sampling_convenience_dump.likelihood.log_likelihood(params)

def log_prior(self, sample):
return _sampling_convenience_dump.priors.ln_prob(
return sampling_convenience_dump.priors.ln_prob(
sample.parameter_only_dict,
normalized=self.normalize_prior,
)
Expand Down Expand Up @@ -1273,8 +1273,8 @@ def step(self):
proposal = self.proposal_cycle.get_proposal()
prop, log_factor = proposal(
self.chain,
likelihood=_sampling_convenience_dump.likelihood,
priors=_sampling_convenience_dump.priors,
likelihood=sampling_convenience_dump.likelihood,
priors=sampling_convenience_dump.priors,
)
logp = self.log_prior(prop)

Expand Down Expand Up @@ -1350,10 +1350,8 @@ def rejection_sample_zero_temperature_samples(self, print_message=False):
zerotemp_logl = hot_samples[LOGLKEY]

# Revert to true likelihood if needed
if _sampling_convenience_dump.use_ratio:
zerotemp_logl += (
_sampling_convenience_dump.likelihood.noise_log_likelihood()
)
if sampling_convenience_dump.use_ratio:
zerotemp_logl += sampling_convenience_dump.likelihood.noise_log_likelihood()

# Calculate normalised weights
log_weights = (1 - beta) * zerotemp_logl
Expand Down Expand Up @@ -1383,9 +1381,9 @@ def get_initial_maximimum_posterior_sample(beta):

"""
logger.info("Finding initial maximum posterior estimate")
likelihood = _sampling_convenience_dump.likelihood
priors = _sampling_convenience_dump.priors
search_parameter_keys = _sampling_convenience_dump.search_parameter_keys
likelihood = sampling_convenience_dump.likelihood
priors = sampling_convenience_dump.priors
search_parameter_keys = sampling_convenience_dump.search_parameter_keys

bounds = []
for key in search_parameter_keys:
Expand All @@ -1398,7 +1396,7 @@ def neg_log_post(x):
if np.isinf(ln_prior):
return -np.inf

parameters = deepcopy(_sampling_convenience_dump.parameters)
parameters = deepcopy(sampling_convenience_dump.parameters)
parameters.update(sample)

return -beta * likelihood.log_likelihood(parameters) - ln_prior
Expand Down
84 changes: 59 additions & 25 deletions bilby/core/result.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,8 @@
from copy import copy
from importlib import import_module
from itertools import product
import multiprocessing
from functools import partial

import numpy as np
import pandas as pd
import scipy.stats
Expand All @@ -33,6 +34,11 @@
EXTENSIONS = ["json", "hdf5", "h5", "pickle", "pkl"]


def __eval_l(likelihood, params):
likelihood.parameters.update(params)
return likelihood.log_likelihood()


def result_file_name(outdir, label, extension='json', gzip=False):
""" Returns the standard filename used for a result file

Expand Down Expand Up @@ -191,7 +197,7 @@ def read_in_result_list(filename_list, invalid="warning"):

def get_weights_for_reweighting(
result, new_likelihood=None, new_prior=None, old_likelihood=None,
old_prior=None, resume_file=None, n_checkpoint=5000, npool=1):
old_prior=None, resume_file=None, n_checkpoint=5000, npool=1, pool=None):
""" Calculate the weights for reweight()

See bilby.core.result.reweight() for help with the inputs
Expand Down Expand Up @@ -234,23 +240,26 @@ def get_weights_for_reweighting(
basedir = os.path.split(resume_file)[0]
check_directory_exists_and_if_not_mkdir(basedir)

dict_samples = [{key: sample[key] for key in result.posterior}
for _, sample in result.posterior.iterrows()]
dict_samples = result.posterior.to_dict(orient="records")
n = len(dict_samples) - starting_index

# Helper function to compute likelihoods in parallel
def eval_pool(this_logl):
with multiprocessing.Pool(processes=npool) as pool:
chunksize = max(100, n // (2 * npool))
return list(tqdm(
pool.imap(
this_logl.log_likelihood,
dict_samples[starting_index:],
chunksize=chunksize,
),
from .utils.parallel import bilby_pool

with bilby_pool(likelihood=this_logl, npool=npool) as my_pool:
if my_pool is None:
map_fn = map
else:
chunksize = max(100, n // (2 * npool))
map_fn = partial(my_pool.imap, chunksize=chunksize)

log_l = list(tqdm(
map_fn(this_logl.log_likelihood, dict_samples[starting_index:]),
desc='Computing likelihoods',
total=n)
)
total=n,
))
return log_l

if old_likelihood is None:
old_log_likelihood_array[starting_index:] = \
Expand Down Expand Up @@ -324,7 +333,7 @@ def rejection_sample(posterior, weights):
def reweight(result, label=None, new_likelihood=None, new_prior=None,
old_likelihood=None, old_prior=None, conversion_function=None, npool=1,
verbose_output=False, resume_file=None, n_checkpoint=5000,
use_nested_samples=False):
use_nested_samples=False, pool=None):
""" Reweight a result to a new likelihood/prior using rejection sampling

Parameters
Expand Down Expand Up @@ -387,7 +396,9 @@ def reweight(result, label=None, new_likelihood=None, new_prior=None,
get_weights_for_reweighting(
result, new_likelihood=new_likelihood, new_prior=new_prior,
old_likelihood=old_likelihood, old_prior=old_prior,
resume_file=resume_file, n_checkpoint=n_checkpoint, npool=npool)
resume_file=resume_file, n_checkpoint=n_checkpoint,
npool=npool, pool=pool,
)

if use_nested_samples:
ln_weights += np.log(result.posterior["weights"])
Expand All @@ -414,10 +425,14 @@ def reweight(result, label=None, new_likelihood=None, new_prior=None,

if conversion_function is not None:
data_frame = result.posterior
if "npool" in inspect.signature(conversion_function).parameters:
data_frame = conversion_function(data_frame, new_likelihood, new_prior, npool=npool)
else:
data_frame = conversion_function(data_frame, new_likelihood, new_prior)
parameters = inspect.signature(conversion_function).parameters
kwargs = dict()
for key, value in [
("likelihood", new_likelihood), ("priors", new_prior), ("npool", npool), ("pool", pool)
]:
if key in parameters:
kwargs[key] = value
data_frame = conversion_function(data_frame, **kwargs)
result.posterior = data_frame

if label:
Expand Down Expand Up @@ -758,6 +773,21 @@ def log_10_evidence_err(self):
def log_10_noise_evidence(self):
return self.log_noise_evidence / np.log(10)

@property
def sampler_kwargs(self):
return self._sampler_kwargs

@sampler_kwargs.setter
def sampler_kwargs(self, sampler_kwargs):
if sampler_kwargs is None:
sampler_kwargs = dict()
else:
sampler_kwargs = copy(sampler_kwargs)
if "pool" in sampler_kwargs:
# pool objects can't be neatly serialized
sampler_kwargs["pool"] = None
self._sampler_kwargs = sampler_kwargs

@property
def version(self):
return self._version
Expand Down Expand Up @@ -1523,7 +1553,7 @@ def _add_prior_fixed_values_to_posterior(posterior, priors):
return posterior

def samples_to_posterior(self, likelihood=None, priors=None,
conversion_function=None, npool=1):
conversion_function=None, npool=1, pool=None):
"""
Convert array of samples to posterior (a Pandas data frame)

Expand Down Expand Up @@ -1553,10 +1583,14 @@ def samples_to_posterior(self, likelihood=None, priors=None,
data_frame['log_prior'] = self.log_prior_evaluations

if conversion_function is not None:
if "npool" in inspect.signature(conversion_function).parameters:
data_frame = conversion_function(data_frame, likelihood, priors, npool=npool)
else:
data_frame = conversion_function(data_frame, likelihood, priors)
parameters = inspect.signature(conversion_function).parameters
kwargs = dict()
for key, value in [
("likelihood", likelihood), ("priors", priors), ("npool", npool), ("pool", pool)
]:
if key in parameters:
kwargs[key] = value
data_frame = conversion_function(data_frame, **kwargs)
self.posterior = data_frame

def calculate_prior_values(self, priors):
Expand Down
Loading
Loading