Skip to content

Commit 04acd21

Browse files
committed
FEAT: improve user pool passing
FEAT: improve reweighting parallelisation FEAT: add parameters as argument to new pool BUG: test that pool exists at cleanup BUG: test pool exists at closing REFACTOR: refactor run_sampler to simplify pool logic DEP: discourage setting up pool in sampler REFACTOR: remove top level multiprocessing import BUG: make sure prior is passed to pool creation BUG: fix test failures TEST: fix reproducibility test BUG: fix a typo in conversion function test MAINT: don't create pool of size 1 BUG: only include chunksize in multiprocessing map DOC: add docstrings for pool functions DOC: update pool docstrings Address review comments TYPO: Fix typo in parameter description comments BUG: move definition of chunk size in reweighting MAINT: remove old function Fix wrong function name in docstring REFACTOR: refactor pool initialization and add documentation
1 parent 8d3df03 commit 04acd21

13 files changed

Lines changed: 743 additions & 333 deletions

File tree

bilby/bilby_mcmc/sampler.py

Lines changed: 20 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -15,7 +15,6 @@
1515
MCMCSampler,
1616
ResumeError,
1717
SamplerError,
18-
_sampling_convenience_dump,
1918
signal_wrapper,
2019
)
2120
from ..core.utils import (
@@ -24,6 +23,7 @@
2423
random,
2524
safe_file_dump,
2625
)
26+
from ..core.utils.parallel import sampling_convenience_dump
2727
from . import proposals
2828
from .chain import Chain, Sample
2929
from .utils import LOGLKEY, LOGPKEY, ConvergenceInputs, ParallelTemperingInputs
@@ -284,7 +284,7 @@ def add_data_to_result(result, ptsampler, outdir, label, make_plots):
284284
total_steps=ptsampler.position,
285285
nsamples=ptsampler.nsamples,
286286
)
287-
if ptsampler.pool is not None:
287+
if ptsampler.pool is not None and hasattr(ptsampler.pool, "_processes"):
288288
npool = ptsampler.pool._processes
289289
else:
290290
npool = 1
@@ -615,7 +615,7 @@ def __init__(
615615

616616
self._nsamples_dict = {}
617617
self.ensemble_proposal_cycle = proposals.get_default_ensemble_proposal_cycle(
618-
_sampling_convenience_dump.priors
618+
sampling_convenience_dump.priors
619619
)
620620
self.sampling_time = 0
621621
self.ln_z_dict = dict()
@@ -630,7 +630,7 @@ def get_initial_betas(self):
630630
elif pt_inputs.Tmax is not None:
631631
betas = np.logspace(0, -np.log10(pt_inputs.Tmax), pt_inputs.ntemps)
632632
elif pt_inputs.Tmax_from_SNR is not None:
633-
ndim = len(_sampling_convenience_dump.priors.non_fixed_keys)
633+
ndim = len(sampling_convenience_dump.priors.non_fixed_keys)
634634
target_hot_likelihood = ndim / 2
635635
Tmax = pt_inputs.Tmax_from_SNR**2 / (2 * target_hot_likelihood)
636636
betas = np.logspace(0, -np.log10(Tmax), pt_inputs.ntemps)
@@ -1172,15 +1172,15 @@ def __init__(
11721172
self.Eindex = Eindex
11731173
self.use_ratio = use_ratio
11741174
self.normalize_prior = normalize_prior
1175-
self.parameters = _sampling_convenience_dump.priors.non_fixed_keys
1175+
self.parameters = sampling_convenience_dump.priors.non_fixed_keys
11761176
self.ndim = len(self.parameters)
11771177

11781178
if initial_sample_method.lower() == "prior":
1179-
full_sample_dict = _sampling_convenience_dump.priors.sample()
1179+
full_sample_dict = sampling_convenience_dump.priors.sample()
11801180
initial_sample = {
11811181
k: v
11821182
for k, v in full_sample_dict.items()
1183-
if k in _sampling_convenience_dump.priors.non_fixed_keys
1183+
if k in sampling_convenience_dump.priors.non_fixed_keys
11841184
}
11851185
elif initial_sample_method.lower() in ["maximize", "maximise", "maximum"]:
11861186
initial_sample = get_initial_maximimum_posterior_sample(self.beta)
@@ -1217,7 +1217,7 @@ def __init__(
12171217

12181218
self.proposal_cycle = proposals.get_proposal_cycle(
12191219
proposal_cycle,
1220-
_sampling_convenience_dump.priors,
1220+
sampling_convenience_dump.priors,
12211221
L1steps=self.chain.L1steps,
12221222
warn=warn,
12231223
)
@@ -1236,16 +1236,16 @@ def set_convergence_inputs(self, convergence_inputs):
12361236
self.stop_after_convergence = convergence_inputs.stop_after_convergence
12371237

12381238
def log_likelihood(self, sample):
1239-
params = deepcopy(_sampling_convenience_dump.parameters)
1239+
params = deepcopy(sampling_convenience_dump.parameters)
12401240
params.update(sample.sample_dict)
12411241

12421242
if self.use_ratio:
1243-
return _sampling_convenience_dump.likelihood.log_likelihood_ratio(params)
1243+
return sampling_convenience_dump.likelihood.log_likelihood_ratio(params)
12441244
else:
1245-
return _sampling_convenience_dump.likelihood.log_likelihood(params)
1245+
return sampling_convenience_dump.likelihood.log_likelihood(params)
12461246

12471247
def log_prior(self, sample):
1248-
return _sampling_convenience_dump.priors.ln_prob(
1248+
return sampling_convenience_dump.priors.ln_prob(
12491249
sample.parameter_only_dict,
12501250
normalized=self.normalize_prior,
12511251
)
@@ -1273,8 +1273,8 @@ def step(self):
12731273
proposal = self.proposal_cycle.get_proposal()
12741274
prop, log_factor = proposal(
12751275
self.chain,
1276-
likelihood=_sampling_convenience_dump.likelihood,
1277-
priors=_sampling_convenience_dump.priors,
1276+
likelihood=sampling_convenience_dump.likelihood,
1277+
priors=sampling_convenience_dump.priors,
12781278
)
12791279
logp = self.log_prior(prop)
12801280

@@ -1350,9 +1350,9 @@ def rejection_sample_zero_temperature_samples(self, print_message=False):
13501350
zerotemp_logl = hot_samples[LOGLKEY]
13511351

13521352
# Revert to true likelihood if needed
1353-
if _sampling_convenience_dump.use_ratio:
1353+
if sampling_convenience_dump.use_ratio:
13541354
zerotemp_logl += (
1355-
_sampling_convenience_dump.likelihood.noise_log_likelihood()
1355+
sampling_convenience_dump.likelihood.noise_log_likelihood()
13561356
)
13571357

13581358
# Calculate normalised weights
@@ -1383,9 +1383,9 @@ def get_initial_maximimum_posterior_sample(beta):
13831383
13841384
"""
13851385
logger.info("Finding initial maximum posterior estimate")
1386-
likelihood = _sampling_convenience_dump.likelihood
1387-
priors = _sampling_convenience_dump.priors
1388-
search_parameter_keys = _sampling_convenience_dump.search_parameter_keys
1386+
likelihood = sampling_convenience_dump.likelihood
1387+
priors = sampling_convenience_dump.priors
1388+
search_parameter_keys = sampling_convenience_dump.search_parameter_keys
13891389

13901390
bounds = []
13911391
for key in search_parameter_keys:
@@ -1398,7 +1398,7 @@ def neg_log_post(x):
13981398
if np.isinf(ln_prior):
13991399
return -np.inf
14001400

1401-
parameters = deepcopy(_sampling_convenience_dump.parameters)
1401+
parameters = deepcopy(sampling_convenience_dump.parameters)
14021402
parameters.update(sample)
14031403

14041404
return -beta * likelihood.log_likelihood(parameters) - ln_prior

bilby/core/result.py

Lines changed: 59 additions & 25 deletions
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,8 @@
88
from copy import copy
99
from importlib import import_module
1010
from itertools import product
11-
import multiprocessing
11+
from functools import partial
12+
1213
import numpy as np
1314
import pandas as pd
1415
import scipy.stats
@@ -33,6 +34,11 @@
3334
EXTENSIONS = ["json", "hdf5", "h5", "pickle", "pkl"]
3435

3536

37+
def __eval_l(likelihood, params):
38+
likelihood.parameters.update(params)
39+
return likelihood.log_likelihood()
40+
41+
3642
def result_file_name(outdir, label, extension='json', gzip=False):
3743
""" Returns the standard filename used for a result file
3844
@@ -191,7 +197,7 @@ def read_in_result_list(filename_list, invalid="warning"):
191197

192198
def get_weights_for_reweighting(
193199
result, new_likelihood=None, new_prior=None, old_likelihood=None,
194-
old_prior=None, resume_file=None, n_checkpoint=5000, npool=1):
200+
old_prior=None, resume_file=None, n_checkpoint=5000, npool=1, pool=None):
195201
""" Calculate the weights for reweight()
196202
197203
See bilby.core.result.reweight() for help with the inputs
@@ -234,23 +240,26 @@ def get_weights_for_reweighting(
234240
basedir = os.path.split(resume_file)[0]
235241
check_directory_exists_and_if_not_mkdir(basedir)
236242

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

241246
# Helper function to compute likelihoods in parallel
242247
def eval_pool(this_logl):
243-
with multiprocessing.Pool(processes=npool) as pool:
244-
chunksize = max(100, n // (2 * npool))
245-
return list(tqdm(
246-
pool.imap(
247-
this_logl.log_likelihood,
248-
dict_samples[starting_index:],
249-
chunksize=chunksize,
250-
),
248+
from .utils.parallel import bilby_pool
249+
250+
with bilby_pool(likelihood=this_logl, npool=npool) as my_pool:
251+
if my_pool is None:
252+
map_fn = map
253+
else:
254+
chunksize = max(100, n // (2 * npool))
255+
map_fn = partial(my_pool.imap, chunksize=chunksize)
256+
257+
log_l = list(tqdm(
258+
map_fn(this_logl.log_likelihood, dict_samples[starting_index:]),
251259
desc='Computing likelihoods',
252-
total=n)
253-
)
260+
total=n,
261+
))
262+
return log_l
254263

255264
if old_likelihood is None:
256265
old_log_likelihood_array[starting_index:] = \
@@ -324,7 +333,7 @@ def rejection_sample(posterior, weights):
324333
def reweight(result, label=None, new_likelihood=None, new_prior=None,
325334
old_likelihood=None, old_prior=None, conversion_function=None, npool=1,
326335
verbose_output=False, resume_file=None, n_checkpoint=5000,
327-
use_nested_samples=False):
336+
use_nested_samples=False, pool=None):
328337
""" Reweight a result to a new likelihood/prior using rejection sampling
329338
330339
Parameters
@@ -387,7 +396,9 @@ def reweight(result, label=None, new_likelihood=None, new_prior=None,
387396
get_weights_for_reweighting(
388397
result, new_likelihood=new_likelihood, new_prior=new_prior,
389398
old_likelihood=old_likelihood, old_prior=old_prior,
390-
resume_file=resume_file, n_checkpoint=n_checkpoint, npool=npool)
399+
resume_file=resume_file, n_checkpoint=n_checkpoint,
400+
npool=npool, pool=pool,
401+
)
391402

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

415426
if conversion_function is not None:
416427
data_frame = result.posterior
417-
if "npool" in inspect.signature(conversion_function).parameters:
418-
data_frame = conversion_function(data_frame, new_likelihood, new_prior, npool=npool)
419-
else:
420-
data_frame = conversion_function(data_frame, new_likelihood, new_prior)
428+
parameters = inspect.signature(conversion_function).parameters
429+
kwargs = dict()
430+
for key, value in [
431+
("likelihood", new_likelihood), ("priors", new_prior), ("npool", npool), ("pool", pool)
432+
]:
433+
if key in parameters:
434+
kwargs[key] = value
435+
data_frame = conversion_function(data_frame, **kwargs)
421436
result.posterior = data_frame
422437

423438
if label:
@@ -758,6 +773,21 @@ def log_10_evidence_err(self):
758773
def log_10_noise_evidence(self):
759774
return self.log_noise_evidence / np.log(10)
760775

776+
@property
777+
def sampler_kwargs(self):
778+
return self._sampler_kwargs
779+
780+
@sampler_kwargs.setter
781+
def sampler_kwargs(self, sampler_kwargs):
782+
if sampler_kwargs is None:
783+
sampler_kwargs = dict()
784+
else:
785+
sampler_kwargs = copy(sampler_kwargs)
786+
if "pool" in sampler_kwargs:
787+
# pool objects can't be neatly serialized
788+
sampler_kwargs["pool"] = None
789+
self._sampler_kwargs = sampler_kwargs
790+
761791
@property
762792
def version(self):
763793
return self._version
@@ -1523,7 +1553,7 @@ def _add_prior_fixed_values_to_posterior(posterior, priors):
15231553
return posterior
15241554

15251555
def samples_to_posterior(self, likelihood=None, priors=None,
1526-
conversion_function=None, npool=1):
1556+
conversion_function=None, npool=1, pool=None):
15271557
"""
15281558
Convert array of samples to posterior (a Pandas data frame)
15291559
@@ -1553,10 +1583,14 @@ def samples_to_posterior(self, likelihood=None, priors=None,
15531583
data_frame['log_prior'] = self.log_prior_evaluations
15541584

15551585
if conversion_function is not None:
1556-
if "npool" in inspect.signature(conversion_function).parameters:
1557-
data_frame = conversion_function(data_frame, likelihood, priors, npool=npool)
1558-
else:
1559-
data_frame = conversion_function(data_frame, likelihood, priors)
1586+
parameters = inspect.signature(conversion_function).parameters
1587+
kwargs = dict()
1588+
for key, value in [
1589+
("likelihood", likelihood), ("priors", priors), ("npool", npool), ("pool", pool)
1590+
]:
1591+
if key in parameters:
1592+
kwargs[key] = value
1593+
data_frame = conversion_function(data_frame, **kwargs)
15601594
self.posterior = data_frame
15611595

15621596
def calculate_prior_values(self, priors):

0 commit comments

Comments
 (0)