From 51b147bc44d47026f6afb622885de05d13d7695b Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=90=B4=E6=B2=81?= Date: Fri, 7 Aug 2026 16:14:21 +0800 Subject: [PATCH 01/13] Add broken power-law luminosity functions --- setup.cfg | 4 +- zdm/MCMC.py | 111 ++++++-- zdm/data/MCMC/1_broken_power_law.json | 81 ++++++ zdm/data/MCMC/2_broken_power_law.json | 90 +++++++ zdm/data/MCMC/gamma.json | 70 +++++ zdm/data/MCMC/power_law.json | 70 +++++ zdm/energetics.py | 357 +++++++++++++++++++++----- zdm/grid.py | 80 +++++- zdm/iteration.py | 48 +++- zdm/parameters.py | 40 ++- zdm/scripts/MCMC/MCMC_wrap.py | 8 +- zdm/tests/test_energetics.py | 166 ++++++++++++ zdm/tests/test_mcmc.py | 128 +++++++++ zdm/tests/test_parameters.py | 15 +- 14 files changed, 1171 insertions(+), 97 deletions(-) create mode 100644 zdm/data/MCMC/1_broken_power_law.json create mode 100644 zdm/data/MCMC/2_broken_power_law.json create mode 100644 zdm/data/MCMC/gamma.json create mode 100644 zdm/data/MCMC/power_law.json diff --git a/setup.cfg b/setup.cfg index 712df842..1c838f14 100644 --- a/setup.cfg +++ b/setup.cfg @@ -46,9 +46,9 @@ install_requires = tqdm>=4.67.1 importlib_resources>=6.0.0 cmasher>=1.9 - ne2001 @ git+https://github.com/FRBs/ne2001 + ne2001 @ git+https://github.com/FRBs/ne2001.git frb @ git+https://github.com/FRBs/FRB - astropath @ git+https://github.com/FRBs/astropath + astropath @ git+https://github.com/FRBs/astropath.git [options.extras_require] test = diff --git a/zdm/MCMC.py b/zdm/MCMC.py index 93bd3620..eb5b18d9 100644 --- a/zdm/MCMC.py +++ b/zdm/MCMC.py @@ -36,7 +36,6 @@ import importlib.resources as resources import emcee -import scipy.stats as st import time from zdm import loading @@ -51,6 +50,59 @@ import os #============================================================================== +def valid_parameter_combination(param_dict, state): + """Check joint constraints that cannot be expressed as 1-D priors. + + For broken power-law luminosity functions, the characteristic energies + are stored in log10 space and must remain strictly ordered. + Values absent from ``param_dict`` are taken from ``state``, allowing any + subset of the energy parameters to be sampled. + """ + luminosity_function = param_dict.get( + 'luminosity_function', state.energy.luminosity_function + ) + lEmin = param_dict.get('lEmin', state.energy.lEmin) + lEmax = param_dict.get('lEmax', state.energy.lEmax) + if luminosity_function == 4: + lEb = param_dict.get('lEb', state.energy.lEb) + return bool(lEmin < lEb < lEmax) + if luminosity_function == 5: + lEb = param_dict.get('lEb', state.energy.lEb) + lEb2 = param_dict.get('lEb2', state.energy.lEb2) + return bool(lEmin < lEb < lEb2 < lEmax) + return True + + +def get_initial_walkers(state, params, nwalkers, rng=None, max_attempts=10000): + """Draw walker positions from the priors, respecting joint constraints.""" + if rng is None: + rng = np.random.default_rng() + + param_names = list(params) + ndim = len(param_names) + walkers = np.empty((nwalkers, ndim), dtype=float) + + for iwalker in range(nwalkers): + for _ in range(max_attempts): + candidate = np.array([ + rng.uniform(params[name]['min'], params[name]['max']) + for name in param_names + ]) + candidate_dict = dict(zip(param_names, candidate)) + if valid_parameter_combination(candidate_dict, state): + walkers[iwalker] = candidate + break + else: + raise ValueError( + "Could not initialize MCMC walkers inside the joint priors. " + "For luminosity_function=4 or 5, ensure the prior ranges " + "permit the required ordering of break energies." + ) + + return walkers + +#============================================================================== + def calc_log_posterior(param_vals, state, params, surveys_sep, Pn=False, pNreps=True, psnr=True, ptauw=False, pwb=False, log_halo=False, lin_host=False, ind_surveys=False, g0info=None): """Calculate log-posterior probability for a parameter vector. @@ -110,6 +162,9 @@ def calc_log_posterior(param_vals, state, params, surveys_sep, Pn=False, pNreps= else: param_dict[key] = param_vals[i] + if in_priors and not valid_parameter_combination(param_dict, state): + in_priors = False + # Initialise list if requesting individual survey likelihoods if ind_surveys: ll_list = [] @@ -163,9 +218,23 @@ def calc_log_posterior(param_vals, state, params, surveys_sep, Pn=False, pNreps= # gets new zDM grid if F and H0 in the param_dict if 'H0' in param_dict or 'logF' in param_dict or g0info is None: datdir = resources.files('zdm').joinpath('GridData') + grid_kwargs = {} + if g0info is not None: + # Preserve the resolution selected by MCMC_wrap. Previously, + # sampling H0/logF silently reverted every worker to the + # 500 x 1400 default grid, causing both shape errors and large + # unexpected memory use in low-resolution pilot runs. + dz = zvals[-1] - zvals[-2] + ddm = dmvals[-1] - dmvals[-2] + grid_kwargs = { + 'nz': zvals.size, + 'zmax': zvals[-1] + dz / 2, + 'ndm': dmvals.size, + 'dmmax': dmvals[-1] + ddm / 2, + } zDMgrid, zvals,dmvals = mf.get_zdm_grid( state, new=True, plot=False, method='analytic', - datdir=datdir) + datdir=datdir, **grid_kwargs) g0info = [zDMgrid, zvals,dmvals] if len(surveys_sep[0]) != 0: @@ -222,7 +291,7 @@ def mcmc_runner(logpf, outfile, state, params, surveys, nwalkers=10, nsteps=100, grid_params (dictionary) = nz, ndm, dmmax nwalkers (int) = Number of walkers nsteps (int) = Number of steps - nthreads (int) = Number of threads (currently not implemented - uses default) + nthreads (int) = Number of worker processes Pn (bool) = Include Pn or not pNreps (bool) = Include pNreps or not ptauw (bool) = Include ptauw or not @@ -236,14 +305,13 @@ def mcmc_runner(logpf, outfile, state, params, surveys, nwalkers=10, nsteps=100, """ ndim = len(params) - starting_guesses = [] - - # Produce starting guesses for each parameter + # Report priors in sampling order. for key,val in params.items(): - starting_guesses.append(st.uniform(loc=val['min'], scale=val['max']-val['min']).rvs(size=[nwalkers])) print(key + " priors: " + str(val['min']) + "," + str(val['max'])) - - starting_guesses = np.array(starting_guesses).T + + # Draw only physically valid initial positions. Broken power-law walkers + # must satisfy the required ordering of their characteristic energies. + starting_guesses = get_initial_walkers(state, params, nwalkers) # we only reset the backend if specifically requested. # This means that walkers will continue from a previous iteration @@ -257,14 +325,16 @@ def mcmc_runner(logpf, outfile, state, params, surveys, nwalkers=10, nsteps=100, start = time.time() - # may or may not be needed - #os.environ["OMP_NUM_THREADS"] = "1" + if nthreads < 1: + raise ValueError("nthreads must be at least 1") + + # Prevent numerical libraries from starting extra threads inside each + # worker process, which can otherwise multiply both CPU and memory use. + os.environ["OMP_NUM_THREADS"] = "1" import multiprocessing as mp Pool = mp.get_context('fork').Pool - - - - with Pool() as pool: # could add mp.Pool(ntrheads=5) or Pool = None + + def run_sampler(pool): sampler = emcee.EnsembleSampler(nwalkers, ndim, logpf, args=[state, params, surveys, Pn, pNreps, psnr, ptauw, pwb, log_halo, lin_host, ind_surveys, g0info], backend=backend, pool=pool) if exists: @@ -273,6 +343,17 @@ def mcmc_runner(logpf, outfile, state, params, surveys, nwalkers=10, nsteps=100, else: # start from new random guesses sampler.run_mcmc(starting_guesses, nsteps, progress=True) + return sampler + + if nthreads == 1: + # Avoid creating a second Python process for the memory-conservative + # default mode. + sampler = run_sampler(None) + else: + # Recycling workers periodically releases arrays retained by Python's + # allocator during repeated six-survey grid construction. + with Pool(processes=nthreads, maxtasksperchild=10) as pool: + sampler = run_sampler(pool) end = time.time() print("Total time taken: " + str(end - start)) diff --git a/zdm/data/MCMC/1_broken_power_law.json b/zdm/data/MCMC/1_broken_power_law.json new file mode 100644 index 00000000..51501af9 --- /dev/null +++ b/zdm/data/MCMC/1_broken_power_law.json @@ -0,0 +1,81 @@ +{ + "mcmc": { + "parameter_order": ["sfr_n", "alpha", "lmean", "lsigma", "lEmin", "lEb", "lEmax", "gamma", "gamma2", "H0"] + }, + "config": { + "luminosity_function": 4, + "alpha_method": 1, + "source_evolution": 0, + "logF": -0.494850021680094, + "DMhalo": 50.0, + "sigmaDMG": 0.2, + "sigmaHalo": 15.0, + "halo_method": 0, + "Wlogmean": 0.0, + "Wlogsigma": 0.42, + "WNbins": 5, + "WidthFunction": 1, + "Wthresh": 0.5, + "Wmethod": 2, + "WMin": 0.1, + "WMax": 100.0, + "WNInternalBins": 100, + "Slogmean": 0.305, + "Slogsigma": 0.75, + "ScatFunction": 1, + "Sbackproject": true, + "Sfnorm": 600.0, + "Sfpower": -4.0, + "Smaxsigma": 3.0 + }, + "sfr_n": { + "DC": "FRBdemo", + "min": -2.0, + "max": 4.0 + }, + "alpha": { + "DC": "energy", + "min": -0.5, + "max": 2.5 + }, + "lmean": { + "DC": "host", + "min": 1.0, + "max": 3.0 + }, + "lsigma": { + "DC": "host", + "min": 0.1, + "max": 1.5 + }, + "lEmin": { + "DC": "energy", + "min": 36.0, + "max": 40.5 + }, + "lEb": { + "DC": "energy", + "min": 36.0, + "max": 45.0 + }, + "lEmax": { + "DC": "energy", + "min": 40.5, + "max": 45.0 + }, + "gamma": { + "DC": "energy", + "min": -3.0, + "max": 0.0 + }, + "gamma2": { + "DC": "energy", + "min": -5.0, + "max": 0.0 + }, + "H0": { + "DC": "cosmo", + "min": 35.0, + "max": 110.0 + } +} diff --git a/zdm/data/MCMC/2_broken_power_law.json b/zdm/data/MCMC/2_broken_power_law.json new file mode 100644 index 00000000..853ae624 --- /dev/null +++ b/zdm/data/MCMC/2_broken_power_law.json @@ -0,0 +1,90 @@ +{ + "mcmc": { + "parameter_order": ["sfr_n", "alpha", "lmean", "lsigma", "lEmin", "lEb", "lEb2", "lEmax", "gamma", "gamma2", "gamma3", "H0"] + }, + "config": { + "luminosity_function": 5, + "alpha_method": 1, + "source_evolution": 0, + "logF": -0.494850021680094, + "DMhalo": 50.0, + "sigmaDMG": 0.2, + "sigmaHalo": 15.0, + "halo_method": 0, + "Wlogmean": 0.0, + "Wlogsigma": 0.42, + "WNbins": 5, + "WidthFunction": 1, + "Wthresh": 0.5, + "Wmethod": 2, + "WMin": 0.1, + "WMax": 100.0, + "WNInternalBins": 100, + "Slogmean": 0.305, + "Slogsigma": 0.75, + "ScatFunction": 1, + "Sfnorm": 600.0, + "Sfpower": -4.0, + "Smaxsigma": 3.0 + }, + "sfr_n": { + "DC": "FRBdemo", + "min": -2.0, + "max": 4.0 + }, + "alpha": { + "DC": "energy", + "min": -0.5, + "max": 2.5 + }, + "lmean": { + "DC": "host", + "min": 1.0, + "max": 3.0 + }, + "lsigma": { + "DC": "host", + "min": 0.1, + "max": 1.5 + }, + "lEmin": { + "DC": "energy", + "min": 36.0, + "max": 40.5 + }, + "lEb": { + "DC": "energy", + "min": 36.0, + "max": 44.0 + }, + "lEb2": { + "DC": "energy", + "min": 37.0, + "max": 45.0 + }, + "lEmax": { + "DC": "energy", + "min": 40.5, + "max": 45.0 + }, + "gamma": { + "DC": "energy", + "min": -3.0, + "max": 0.0 + }, + "gamma2": { + "DC": "energy", + "min": -5.0, + "max": 0.0 + }, + "gamma3": { + "DC": "energy", + "min": -7.0, + "max": 0.0 + }, + "H0": { + "DC": "cosmo", + "min": 35.0, + "max": 110.0 + } +} diff --git a/zdm/data/MCMC/gamma.json b/zdm/data/MCMC/gamma.json new file mode 100644 index 00000000..ccca63f0 --- /dev/null +++ b/zdm/data/MCMC/gamma.json @@ -0,0 +1,70 @@ +{ + "mcmc": { + "parameter_order": ["sfr_n", "alpha", "lmean", "lsigma", "lEmin", "lEmax", "gamma", "H0"] + }, + "config": { + "luminosity_function": 2, + "alpha_method": 1, + "source_evolution": 0, + "logF": -0.494850021680094, + "DMhalo": 50.0, + "sigmaDMG": 0.2, + "sigmaHalo": 15.0, + "halo_method": 0, + "Wlogmean": 0.0, + "Wlogsigma": 0.42, + "WNbins": 5, + "WidthFunction": 1, + "Wthresh": 0.5, + "Wmethod": 2, + "WMin": 0.1, + "WMax": 100.0, + "WNInternalBins": 100, + "Slogmean": 0.305, + "Slogsigma": 0.75, + "ScatFunction": 1, + "Sfnorm": 600.0, + "Sfpower": -4.0, + "Smaxsigma": 3.0 + }, + "sfr_n": { + "DC": "FRBdemo", + "min": -2.0, + "max": 4.0 + }, + "alpha": { + "DC": "energy", + "min": -0.5, + "max": 2.5 + }, + "lmean": { + "DC": "host", + "min": 1.0, + "max": 3.0 + }, + "lsigma": { + "DC": "host", + "min": 0.1, + "max": 1.5 + }, + "lEmin": { + "DC": "energy", + "min": 36.0, + "max": 40.5 + }, + "lEmax": { + "DC": "energy", + "min": 40.5, + "max": 45.0 + }, + "gamma": { + "DC": "energy", + "min": -3.0, + "max": 0.0 + }, + "H0": { + "DC": "cosmo", + "min": 35.0, + "max": 110.0 + } +} diff --git a/zdm/data/MCMC/power_law.json b/zdm/data/MCMC/power_law.json new file mode 100644 index 00000000..27fefaeb --- /dev/null +++ b/zdm/data/MCMC/power_law.json @@ -0,0 +1,70 @@ +{ + "mcmc": { + "parameter_order": ["sfr_n", "alpha", "lmean", "lsigma", "lEmin", "lEmax", "gamma", "H0"] + }, + "config": { + "luminosity_function": 0, + "alpha_method": 1, + "source_evolution": 0, + "logF": -0.494850021680094, + "DMhalo": 50.0, + "sigmaDMG": 0.2, + "sigmaHalo": 15.0, + "halo_method": 0, + "Wlogmean": 0.0, + "Wlogsigma": 0.42, + "WNbins": 5, + "WidthFunction": 1, + "Wthresh": 0.5, + "Wmethod": 2, + "WMin": 0.1, + "WMax": 100.0, + "WNInternalBins": 100, + "Slogmean": 0.305, + "Slogsigma": 0.75, + "ScatFunction": 1, + "Sfnorm": 600.0, + "Sfpower": -4.0, + "Smaxsigma": 3.0 + }, + "sfr_n": { + "DC": "FRBdemo", + "min": -2.0, + "max": 4.0 + }, + "alpha": { + "DC": "energy", + "min": -0.5, + "max": 2.5 + }, + "lmean": { + "DC": "host", + "min": 1.0, + "max": 3.0 + }, + "lsigma": { + "DC": "host", + "min": 0.1, + "max": 1.5 + }, + "lEmin": { + "DC": "energy", + "min": 36.0, + "max": 41.0 + }, + "lEmax": { + "DC": "energy", + "min": 41.0, + "max": 45.0 + }, + "gamma": { + "DC": "energy", + "min": -3.0, + "max": 0.0 + }, + "H0": { + "DC": "cosmo", + "min": 35.0, + "max": 110.0 + } +} diff --git a/zdm/energetics.py b/zdm/energetics.py index f47b0b8a..f3fce691 100644 --- a/zdm/energetics.py +++ b/zdm/energetics.py @@ -15,6 +15,8 @@ Key Functions ------------- - `vector_cum_power_law`: Cumulative power-law luminosity function +- `vector_cum_broken_power_law`: Cumulative broken power-law luminosity function +- `vector_cum_double_broken_power_law`: Cumulative two-break power-law luminosity function - `vector_cum_gamma_spline`: Cumulative gamma function with spline interpolation - `array_cum_gamma_spline`: N-dimensional array wrapper for gamma function - `init_igamma_splines`: Initialize spline lookup tables for gamma functions @@ -94,8 +96,8 @@ def init_igamma_splines(gammas, reinit=False, k=3): igamma_splines[gamma] = interpolate.splrep(lavals, lnumer,k=k) else: igamma_splines[gamma] = interpolate.splrep(avals, numer,k=k) - - + + def init_igamma_linear(gammas: list, reinit: bool = False, log: bool = False): """Initialize linear interpolators for the upper incomplete gamma function. @@ -159,7 +161,42 @@ def template_vector_cumulative_luminosity_function(Eth,*params): return None ########### simple power law functions ############# - + +def array_diff_power_law(Eth,*params): + """ Calculates the differential fraction of bursts for a power law + at a given Eth, where Eth is an N-dimensional array + """ + dims=Eth.shape + Eth=Eth.flatten() + #if gamma >= 0: #handles crazy dodgy cases. Or just return 0? + # result=np.zeros([Eth.size]) + # result[np.where(Eth < Emax)]=1. + # result=result.reshape(dims) + # Eth=Eth.reshape(dims) + # return result + + result=vector_diff_power_law(Eth,*params) + result=result.reshape(dims) + return result + + +def vector_diff_power_law(Eth,*params): + Emin=params[0] + Emax=params[1] + gamma=params[2] + + result=-(gamma*Eth**(gamma-1)) / (Emin**gamma-Emax**gamma ) + + low=np.where(Eth < Emin)[0] + if len(low) > 0: + result[low]=0. + high=np.where(Eth > Emax)[0] + if len(high) > 0: + result[high]=0. + + return result + + def array_cum_power_law(Eth,*params): """ Calculates the fraction of bursts above a certain power law for a given Eth, where Eth is an N-dimensional array @@ -176,10 +213,6 @@ def array_cum_power_law(Eth,*params): result=result.reshape(dims) return result -############## this section defines different luminosity functions ########## - -########### simple power law functions ############# - def vector_cum_power_law(Eth, *params): """Cumulative power-law luminosity function. @@ -212,57 +245,265 @@ def vector_cum_power_law(Eth, *params): result[high]=0. return result -def array_diff_power_law(Eth,*params): - """ Calculates the differential fraction of bursts for a power law - at a given Eth, where Eth is an N-dimensional array + +########### simple broken power law functions ############# + +def vector_cum_broken_power_law(Eth, *params): + """Cumulative broken power-law luminosity function. + + Computes the fraction of bursts with energy above ``Eth`` for a broken + power law whose indices below and above the break energy are ``gamma1`` + and ``gamma2``, respectively. The two branches have the same + normalization at the break. + + Parameters + ---------- + Eth : ndarray + One-dimensional array of energy threshold values. + *params : tuple + ``(Emin, Emax, gamma1, gamma2, Eb)``: minimum energy, maximum + energy, lower power-law index, upper power-law index, and break + energy. + + Returns + ------- + ndarray + Fraction of bursts with ``E > Eth``. The result is 1 below + ``Emin`` and 0 above ``Emax``. + + Notes + ----- + Zero-valued indices are evaluated using their logarithmic limits. """ - dims=Eth.shape - Eth=Eth.flatten() - #if gamma >= 0: #handles crazy dodgy cases. Or just return 0? - # result=np.zeros([Eth.size]) - # result[np.where(Eth < Emax)]=1. - # result=result.reshape(dims) - # Eth=Eth.reshape(dims) - # return result - - result=vector_diff_power_law(Eth,*params) - result=result.reshape(dims) + Emin, Emax, gamma1, gamma2, Eb = params + + if not (0 < Emin < Eb < Emax): + raise ValueError("energies must satisfy 0 < Emin < Eb < Emax") + + Eth = np.asarray(Eth, dtype=float) + scalar_input = Eth.ndim == 0 + if Eth.ndim > 1: + raise ValueError("Eth must be a one-dimensional array") + Eth = np.atleast_1d(Eth) + + def lower_integral(ratio, gamma): + """Return (1 - ratio**gamma) / gamma stably.""" + if gamma == 0: + return -np.log(ratio) + return -np.expm1(gamma * np.log(ratio)) / gamma + + def upper_integral(ratio, gamma): + """Return (ratio**gamma - 1) / gamma stably.""" + if gamma == 0: + return np.log(ratio) + return np.expm1(gamma * np.log(ratio)) / gamma + + upper_at_break = upper_integral(Emax / Eb, gamma2) + normalization = (lower_integral(Emin / Eb, gamma1) + upper_at_break) + + result = np.empty_like(Eth) + below_min = Eth < Emin + above_max = Eth > Emax + below_break = (Eth >= Emin) & (Eth < Eb) + above_break = (Eth >= Eb) & (Eth <= Emax) + + result[below_min] = 1.0 + result[above_max] = 0.0 + result[below_break] = (lower_integral(Eth[below_break] / Eb, gamma1) + upper_at_break) / normalization + result[above_break] = upper_integral(Emax / Eb, gamma2) - upper_integral(Eth[above_break] / Eb, gamma2) + result[above_break] /= normalization + + if scalar_input: + return result[0] return result - -def array_cum_power_law(Eth,*params): - """ Calculates the fraction of bursts above a certain power law - for a given Eth, where Eth is an N-dimensional array + +def array_cum_broken_power_law(Eth, *params): + """N-dimensional wrapper for :func:`vector_cum_broken_power_law`.""" + Eth = np.asarray(Eth) + dims = Eth.shape + result = vector_cum_broken_power_law(Eth.flatten(), *params) + return result.reshape(dims) + + +def vector_diff_broken_power_law(Eth, *params): + """Differential broken power-law luminosity function. + + This is the normalized probability density corresponding to + :func:`vector_cum_broken_power_law`. + + Parameters + ---------- + Eth : ndarray + One-dimensional array of energies. + *params : tuple + ``(Emin, Emax, gamma1, gamma2, Eb)``. + + Returns + ------- + ndarray + Probability density at each energy, with zero density outside + ``[Emin, Emax]``. """ - dims=Eth.shape - Eth=Eth.flatten() - #if gamma >= 0: #handles crazy dodgy cases. Or just return 0? - # result=np.zeros([Eth.size]) - # result[np.where(Eth < Emax)]=1. - # result=result.reshape(dims) - # Eth=Eth.reshape(dims) - # return result - result=vector_cum_power_law(Eth,*params) - result=result.reshape(dims) + Emin, Emax, gamma1, gamma2, Eb = params + + if not (0 < Emin < Eb < Emax): + raise ValueError("energies must satisfy 0 < Emin < Eb < Emax") + + Eth = np.asarray(Eth, dtype=float) + scalar_input = Eth.ndim == 0 + if Eth.ndim > 1: + raise ValueError("Eth must be a one-dimensional array") + Eth = np.atleast_1d(Eth) + + def integral(ratio, gamma): + if gamma == 0: + return np.log(ratio) + return np.expm1(gamma * np.log(ratio)) / gamma + + normalization = (-integral(Emin / Eb, gamma1) + integral(Emax / Eb, gamma2)) + + result = np.zeros_like(Eth) + lower = (Eth >= Emin) & (Eth < Eb) + upper = (Eth >= Eb) & (Eth <= Emax) + result[lower] = ((Eth[lower] / Eb) ** (gamma1 - 1) / (Eb * normalization)) + result[upper] = ((Eth[upper] / Eb) ** (gamma2 - 1) / (Eb * normalization)) + if scalar_input: + return result[0] return result -def vector_diff_power_law(Eth,*params): - Emin=params[0] - Emax=params[1] - gamma=params[2] - - result=-(gamma*Eth**(gamma-1)) / (Emin**gamma-Emax**gamma ) - - low=np.where(Eth < Emin)[0] - if len(low) > 0: - result[low]=0. - high=np.where(Eth > Emax)[0] - if len(high) > 0: - result[high]=0. - + +def array_diff_broken_power_law(Eth, *params): + """N-dimensional wrapper for :func:`vector_diff_broken_power_law`.""" + Eth = np.asarray(Eth) + dims = Eth.shape + result = vector_diff_broken_power_law(Eth.flatten(), *params) + return result.reshape(dims) + + + +########### Double broken power law functions ############# + + +def vector_cum_double_broken_power_law(Eth, *params): + """Cumulative doubly-broken power-law luminosity function. + + The three differential branches are continuous at both break energies + and follow the zDM convention ``dP/dE ~ E**(gamma_i - 1)``. + + Parameters + ---------- + Eth : float or ndarray + Energy threshold values. + *params : tuple + ``(Emin, Emax, gamma1, gamma2, gamma3, Eb1, Eb2)``. + + Returns + ------- + float or ndarray + Fraction of bursts with energy greater than ``Eth``. + """ + Emin, Emax, gamma1, gamma2, gamma3, Eb1, Eb2 = params + + if not (0 < Emin < Eb1 < Eb2 < Emax): + raise ValueError("energies must satisfy 0 < Emin < Eb1 < Eb2 < Emax") + + Eth = np.asarray(Eth, dtype=float) + scalar_input = Eth.ndim == 0 + if Eth.ndim > 1: + raise ValueError("Eth must be a one-dimensional array") + Eth = np.atleast_1d(Eth) + + def lower_integral(ratio, gamma): + """Return (1 - ratio**gamma) / gamma stably.""" + if gamma == 0: + return -np.log(ratio) + return -np.expm1(gamma * np.log(ratio)) / gamma + + def upper_integral(ratio, gamma): + """Return (ratio**gamma - 1) / gamma stably.""" + if gamma == 0: + return np.log(ratio) + return np.expm1(gamma * np.log(ratio)) / gamma + + break_ratio = Eb2 / Eb1 + middle_total = upper_integral(break_ratio, gamma2) + # The extra factor follows from continuity of dP/dE at Eb2. + high_scale = break_ratio ** gamma2 + high_total = high_scale * upper_integral(Emax / Eb2, gamma3) + normalization = (lower_integral(Emin / Eb1, gamma1) + middle_total + high_total) + + result = np.empty_like(Eth) + below_min = Eth < Emin + above_max = Eth > Emax + first = (Eth >= Emin) & (Eth < Eb1) + second = (Eth >= Eb1) & (Eth < Eb2) + third = (Eth >= Eb2) & (Eth <= Emax) + + result[below_min] = 1.0 + result[above_max] = 0.0 + result[first] = (lower_integral(Eth[first] / Eb1, gamma1) + middle_total + high_total) / normalization + result[second] = (middle_total - upper_integral(Eth[second] / Eb1, gamma2) + high_total) / normalization + result[third] = high_scale * (upper_integral(Emax / Eb2, gamma3) - upper_integral(Eth[third] / Eb2, gamma3)) / normalization + + if scalar_input: + return result[0] + return result + + +def array_cum_double_broken_power_law(Eth, *params): + """N-dimensional wrapper for the cumulative doubly-broken power law.""" + Eth = np.asarray(Eth) + dims = Eth.shape + result = vector_cum_double_broken_power_law(Eth.flatten(), *params) + return result.reshape(dims) + + +def vector_diff_double_broken_power_law(Eth, *params): + """Differential probability density for the doubly-broken power law.""" + Emin, Emax, gamma1, gamma2, gamma3, Eb1, Eb2 = params + + if not (0 < Emin < Eb1 < Eb2 < Emax): + raise ValueError("energies must satisfy 0 < Emin < Eb1 < Eb2 < Emax") + + Eth = np.asarray(Eth, dtype=float) + scalar_input = Eth.ndim == 0 + if Eth.ndim > 1: + raise ValueError("Eth must be a one-dimensional array") + Eth = np.atleast_1d(Eth) + + def integral(ratio, gamma): + if gamma == 0: + return np.log(ratio) + return np.expm1(gamma * np.log(ratio)) / gamma + + break_ratio = Eb2 / Eb1 + high_scale = break_ratio ** gamma2 + normalization = (-integral(Emin / Eb1, gamma1)+ integral(break_ratio, gamma2)+ high_scale * integral(Emax / Eb2, gamma3)) + + result = np.zeros_like(Eth) + first = (Eth >= Emin) & (Eth < Eb1) + second = (Eth >= Eb1) & (Eth < Eb2) + third = (Eth >= Eb2) & (Eth <= Emax) + + result[first] = ((Eth[first] / Eb1) ** (gamma1 - 1) / (Eb1 * normalization)) + result[second] = ((Eth[second] / Eb1) ** (gamma2 - 1) / (Eb1 * normalization)) + result[third] = (break_ratio ** (gamma2 - 1) * (Eth[third] / Eb2) ** (gamma3 - 1) / (Eb1 * normalization)) + + if scalar_input: + return result[0] return result +def array_diff_double_broken_power_law(Eth, *params): + """N-dimensional wrapper for the differential doubly-broken power law.""" + Eth = np.asarray(Eth) + dims = Eth.shape + result = vector_diff_double_broken_power_law(Eth.flatten(), *params) + return result.reshape(dims) + + + ########### gamma functions ############# def vector_cum_gamma(Eth, *params): @@ -329,7 +570,7 @@ def vector_cum_gamma_spline(Eth: np.ndarray, *params): Automatically initializes splines for new gamma values if needed. """ global SplineLog - + params=np.array(params) Emin=params[0] Emax=params[1] @@ -348,7 +589,7 @@ def vector_cum_gamma_spline(Eth: np.ndarray, *params): # Low end low = Eth < Emin - + if np.isscalar(result): if low: result = 1. @@ -373,7 +614,7 @@ def vector_cum_gamma_linear(Eth:np.ndarray, *params): # Calculate norm = float(mpmath.gammainc(gamma, a=Emin/Emax)) - + # Branch either with log10 space or without if log: Eth_Emax = Eth - np.log10(Emax) @@ -390,10 +631,10 @@ def vector_cum_gamma_linear(Eth:np.ndarray, *params): Eth_Emax = Eth/Emax if gamma not in igamma_linear.keys(): init_igamma_linear([gamma], log=log) - + numer = igamma_linear[gamma](Eth_Emax) Emin_temp = Emin - + result=numer/norm # Low end @@ -458,11 +699,11 @@ def vector_diff_gamma(Eth, *params): Emin=params[0] Emax=params[1] gamma=params[2] - + norm = Emax*float(mpmath.gammainc(gamma, a=Emin/Emax)) result= (Eth/Emax)**(gamma-1) * np.exp(-Eth/Emax) / norm - + low= Eth < Emin - result[low]=0. - + result[low]=0. + return result diff --git a/zdm/grid.py b/zdm/grid.py index 5e940546..e41847dc 100644 --- a/zdm/grid.py +++ b/zdm/grid.py @@ -168,11 +168,84 @@ def init_luminosity_functions(self): self.vector_cum_lf = energetics.vector_cum_gamma_linear self.array_diff_lf = energetics.array_diff_gamma self.vector_diff_lf = energetics.vector_diff_gamma + elif self.luminosity_function == 4: # Broken power law + self.array_cum_lf = self._array_cum_broken_power_law + self.vector_cum_lf = self._vector_cum_broken_power_law + self.array_diff_lf = self._array_diff_broken_power_law + self.vector_diff_lf = self._vector_diff_broken_power_law + elif self.luminosity_function == 5: # Two-broken power law + self.array_cum_lf = self._array_cum_double_broken_power_law + self.vector_cum_lf = self._vector_cum_double_broken_power_law + self.array_diff_lf = self._array_diff_double_broken_power_law + self.vector_diff_lf = self._vector_diff_double_broken_power_law else: raise ValueError( - "Luminosity function must be 0, not ", self.luminosity_function + "Luminosity function must be one of 0, 1, 2, 3, 4, or 5; " + f"got {self.luminosity_function}" ) + def _broken_power_law_params(self, Emin, Emax, gamma1): + """Expand the standard grid LF arguments for a broken power law.""" + return ( + Emin, + Emax, + gamma1, + self.state.energy.gamma2, + 10 ** self.state.energy.lEb, + ) + + def _array_cum_broken_power_law(self, Eth, Emin, Emax, gamma1, *unused): + params = self._broken_power_law_params(Emin, Emax, gamma1) + return energetics.array_cum_broken_power_law(Eth, *params) + + def _vector_cum_broken_power_law(self, Eth, Emin, Emax, gamma1, *unused): + params = self._broken_power_law_params(Emin, Emax, gamma1) + return energetics.vector_cum_broken_power_law(Eth, *params) + + def _array_diff_broken_power_law(self, Eth, Emin, Emax, gamma1, *unused): + params = self._broken_power_law_params(Emin, Emax, gamma1) + return energetics.array_diff_broken_power_law(Eth, *params) + + def _vector_diff_broken_power_law(self, Eth, Emin, Emax, gamma1, *unused): + params = self._broken_power_law_params(Emin, Emax, gamma1) + return energetics.vector_diff_broken_power_law(Eth, *params) + + def _double_broken_power_law_params(self, Emin, Emax, gamma1): + """Expand the standard grid arguments for a two-break power law.""" + return ( + Emin, + Emax, + gamma1, + self.state.energy.gamma2, + self.state.energy.gamma3, + 10 ** self.state.energy.lEb, + 10 ** self.state.energy.lEb2, + ) + + def _array_cum_double_broken_power_law( + self, Eth, Emin, Emax, gamma1, *unused + ): + params = self._double_broken_power_law_params(Emin, Emax, gamma1) + return energetics.array_cum_double_broken_power_law(Eth, *params) + + def _vector_cum_double_broken_power_law( + self, Eth, Emin, Emax, gamma1, *unused + ): + params = self._double_broken_power_law_params(Emin, Emax, gamma1) + return energetics.vector_cum_double_broken_power_law(Eth, *params) + + def _array_diff_double_broken_power_law( + self, Eth, Emin, Emax, gamma1, *unused + ): + params = self._double_broken_power_law_params(Emin, Emax, gamma1) + return energetics.array_diff_double_broken_power_law(Eth, *params) + + def _vector_diff_double_broken_power_law( + self, Eth, Emin, Emax, gamma1, *unused + ): + params = self._double_broken_power_law_params(Emin, Emax, gamma1) + return energetics.vector_diff_double_broken_power_law(Eth, *params) + def parse_grid(self, zDMgrid, zvals, dmvals): self.grid = zDMgrid self.zvals = zvals @@ -1141,6 +1214,10 @@ def update(self, vparams: dict, ALL=False, prev_grid=None): self.chk_upd_param("lEmin", vparams, update=True), self.chk_upd_param("lEmax", vparams, update=True), self.chk_upd_param("gamma", vparams, update=True), + self.chk_upd_param("gamma2", vparams, update=True), + self.chk_upd_param("gamma3", vparams, update=True), + self.chk_upd_param("lEb", vparams, update=True), + self.chk_upd_param("lEb2", vparams, update=True), ] ): calc_pdv = True @@ -1348,4 +1425,3 @@ def construct_fz(self,ffile,zfile): #np.save(path+"/"+name+"_fz",newfz) #np.save(path+"/"+name+"_z",newz) - diff --git a/zdm/iteration.py b/zdm/iteration.py index 78534ed7..f242588d 100644 --- a/zdm/iteration.py +++ b/zdm/iteration.py @@ -37,6 +37,7 @@ from zdm import cosmology as cos from scipy.stats import poisson import scipy.stats as st +from zdm import energetics from zdm import repeat_grid as zdm_repeat_grid @@ -350,7 +351,7 @@ def calc_likelihoods_1D(grid,survey,doplot=False,norm=True,pdmz=True,psnr=True, ztDMobs=survey.DMEGs[noztaulist] # gets indices of noztaulist within nozlist - tz_tomult = tomult[:,:inoztaulist] + tz_tomult = tomult[:, inoztaulist] # This could all be precalculated within the survey. iws1,iws2,dkws1,dkws2 = survey.get_w_coeffs(Wobs) # total width in survey width bins @@ -369,8 +370,8 @@ def calc_likelihoods_1D(grid,survey,doplot=False,norm=True,pdmz=True,psnr=True, + survey.ptaus[:,itaus2,iws2]*dktaus1*dkws2 # we now multiply by the z-dependencies - ptaus *= zt_tomult - piws *= zt_tomult + ptaus *= tz_tomult + piws *= tz_tomult # sum down the redshift axis to get sum p(tau,w|z)*p(z) ptaus = np.sum(ptaus,axis=0) @@ -534,7 +535,7 @@ def calc_likelihoods_1D(grid,survey,doplot=False,norm=True,pdmz=True,psnr=True, OK = np.where(cumulative > 0)[0] if zwidths: - usew = usew[OK] + usew = usew.flatten()[OK] psnr_gbws[i,j,OK] = differential[OK]/cumulative[OK] @@ -1307,9 +1308,9 @@ def calc_likelihoods_2D(grid,survey,doplot=False,norm=True,pdmz=True,psnr=True,p # multiplies by the width and beam weights for that FRB. These are pre-calculated in the survey # each component below is a vector over nfrb - psnrbw += psnrbws[i,j,:]*zbweights[:,i]*zwweights[:,j] - psnr_gbw += psnr_gbws[i,j,:] *zbweights[:,i]*zwweights[:,j] - pbw += pbws[i,j,:]*zbweights[:,i]*zwweights[:,j] + psnrbw += psnrbws[i,j,:]*bweights[:,i]*wweights[:,j] + psnr_gbw += psnr_gbws[i,j,:] *bweights[:,i]*wweights[:,j] + pbw += pbws[i,j,:]*bweights[:,i]*wweights[:,j] # normalises pbw by normalised sum over all b,w. This gives dual p(b,w) for each FRB @@ -1317,19 +1318,19 @@ def calc_likelihoods_2D(grid,survey,doplot=False,norm=True,pdmz=True,psnr=True,p psnrbw = psnrbw / pwb_norm # psnr_gbws needs no normalisation, provided weights in each dimension sum to unity. But we check here just to be sure - psnr_gbw = psnr_gbw / (np.sum(zbweights,axis=1) * np.sum(zwweights,axis=1)) - psnrbw = psnrbw / (np.sum(zbweights,axis=1) * np.sum(zwweights,axis=1)) + psnr_gbw = psnr_gbw / (np.sum(bweights,axis=1) * np.sum(wweights,axis=1)) + psnrbw = psnrbw / (np.sum(bweights,axis=1) * np.sum(wweights,axis=1)) # calculates p(w) values # then normalises probability over all pbw for j,w in enumerate(grid.eff_weights): - pw[:] += pw_norm[j,:]*zwweights[:,j] + pw[:] += pw_norm[j,:]*wweights[:,j] pw = pw/pwb_norm # calculates p(b) values. # then normalised probability over all pbw for i,b in enumerate(survey.beam_b): - pb[:] += pb_norm[i,:]*zbweights[:,i] + pb[:] += pb_norm[i,:]*bweights[:,i] pb = pb/pwb_norm # calculates p(b|w,z,dM), using p(b|w) p(w) = p(b,w) @@ -1517,8 +1518,27 @@ def ConvertToMeaningfulConstant(state,Eref=1e39): gamma=state.energy.gamma if state.energy.luminosity_function == 0: factor=(Eref/Emin)**gamma - (Emax/Emin)**gamma + elif state.energy.luminosity_function == 4: + factor = energetics.vector_cum_broken_power_law( + np.array([Eref]), + Emin, + Emax, + gamma, + state.energy.gamma2, + 10 ** state.energy.lEb, + ) + elif state.energy.luminosity_function == 5: + factor = energetics.vector_cum_double_broken_power_law( + np.array([Eref]), + Emin, + Emax, + gamma, + state.energy.gamma2, + state.energy.gamma3, + 10 ** state.energy.lEb, + 10 ** state.energy.lEb2, + ) else: - from zdm import energetics factor = energetics.vector_cum_gamma(np.array([Eref]),Emin,Emax,gamma) const *= factor return const @@ -1683,7 +1703,9 @@ def minimise_const_only(vparams:dict,grids:list,surveys:list, else: result=minimize(minus_poisson_ps,startlog10C, args=data,bounds=bounds) - dC=result.x + # scipy returns a one-element array for this one-dimensional + # optimization. Convert explicitly for NumPy 2 compatibility. + dC = result.x.item() t1=time.process_time() # constant needs to include the starting value of .lC diff --git a/zdm/parameters.py b/zdm/parameters.py index 37f162c1..ba01c457 100644 --- a/zdm/parameters.py +++ b/zdm/parameters.py @@ -417,7 +417,7 @@ class EnergeticsParams(data_class.myDataClass): }, ) lEmax: float = field( - default=41.84, + default=43.0, metadata={ "help": "$\log_{10}$ of maximum FRB energy", "unit": "erg", @@ -435,15 +435,47 @@ class EnergeticsParams(data_class.myDataClass): gamma: float = field( default=-1.16, metadata={ - "help": "slope of luminosity distribution function", + "help": "slope of luminosity distribution function; gamma1 for a broken power law", "unit": "", "Notation": "\gamma", }, ) + gamma2: float = field( + default=-2.0, + metadata={ + "help": "second slope of a broken power-law luminosity function", + "unit": "", + "Notation": "\\gamma_2", + }, + ) + gamma3: float = field( + default=-3.0, + metadata={ + "help": "highest-energy slope of a doubly-broken power-law luminosity function", + "unit": "", + "Notation": "\\gamma_3", + }, + ) + lEb: float = field( + default=40.0, + metadata={ + "help": "$\\log_{10}$ of the break energy for a broken power law", + "unit": "erg", + "Notation": "\\log_{10} E_{\\rm b}", + }, + ) + lEb2: float = field( + default=41.0, + metadata={ + "help": "$\\log_{10}$ of the second break energy for a doubly-broken power law", + "unit": "erg", + "Notation": "\\log_{10} E_{\\rm b,2}", + }, + ) luminosity_function: int = field( default=2, metadata={ - "help": "luminosity function applied (0=power-law, 1=gamma, 2=spline+gamma, 3=gamma+linear+log10)" + "help": "luminosity function applied (0=power-law, 1=gamma, 2=spline+gamma, 3=gamma+linear+log10, 4=broken power-law, 5=doubly-broken power-law)" }, ) @@ -553,5 +585,3 @@ class PhotometricParams(data_class.myDataClass): sigma:float =field(default=0.035) sigma_width:int =field(default=6) - - diff --git a/zdm/scripts/MCMC/MCMC_wrap.py b/zdm/scripts/MCMC/MCMC_wrap.py index 60fa25bc..a24e4899 100644 --- a/zdm/scripts/MCMC/MCMC_wrap.py +++ b/zdm/scripts/MCMC/MCMC_wrap.py @@ -87,7 +87,13 @@ def main(): state.set_astropy_cosmo(Planck18) state.update_params(mcmc_dict["config"]) + if args.ptauw: + state.scat.Sbackproject = True + state.width.Wmethod = 3 + print("Config: ", mcmc_dict["config"]) + if args.ptauw: + print("ptauw initialisation: Sbackproject=True, Wmethod=3") print('Pn:', args.Pn) print('psnr:', args.psnr) @@ -111,7 +117,7 @@ def main(): g0info = [zDMgrid,zvals,dmvals] # set z-dependent weights in surveys - if ('Wlogmean' in params or 'Wlogsigma' in params or \ + if (args.ptauw or 'Wlogmean' in params or 'Wlogsigma' in params or \ 'Slogmean' in params or 'Slogsigma' in params): survey_dict = {"WMETHOD": 3} else: diff --git a/zdm/tests/test_energetics.py b/zdm/tests/test_energetics.py index 7bf39ccf..853326c5 100644 --- a/zdm/tests/test_energetics.py +++ b/zdm/tests/test_energetics.py @@ -21,3 +21,169 @@ def test_init_gamma(): assert np.isclose(float(energetics.igamma_linear_log10[-1](0.)), float(energetics.igamma_linear[-1](1.)), rtol=1e-3) + + +def test_vector_cum_broken_power_law(): + """The broken power law follows the requested piecewise expression.""" + Emin, Emax = 1.0, 100.0 + gamma1, gamma2, Eb = -1.0, -2.0, 10.0 + Eth = np.array([0.5, Emin, 5.0, Eb, 50.0, Emax, 200.0]) + + result = energetics.vector_cum_broken_power_law( + Eth, Emin, Emax, gamma1, gamma2, Eb + ) + + denominator = ( + (1 - (Emin / Eb) ** gamma1) / gamma1 + + ((Emax / Eb) ** gamma2 - 1) / gamma2 + ) + expected_below_break = ( + (1 - (5.0 / Eb) ** gamma1) / gamma1 + + ((Emax / Eb) ** gamma2 - 1) / gamma2 + ) / denominator + expected_at_break = ( + ((Emax / Eb) ** gamma2 - 1) / gamma2 + ) / denominator + expected_above_break = ( + ((Emax / Eb) ** gamma2 - (50.0 / Eb) ** gamma2) / gamma2 + ) / denominator + + expected = np.array([ + 1.0, + 1.0, + expected_below_break, + expected_at_break, + expected_above_break, + 0.0, + 0.0, + ]) + np.testing.assert_allclose(result, expected) + assert np.all(np.diff(result) <= 0) + + +def test_vector_cum_broken_power_law_zero_indices(): + """Zero indices use the finite logarithmic limit of the expression.""" + result = energetics.vector_cum_broken_power_law( + np.array([1.0, 10.0, 100.0]), 1.0, 100.0, 0.0, 0.0, 10.0 + ) + + np.testing.assert_allclose(result, [1.0, 0.5, 0.0]) + + +def test_diff_broken_power_law_matches_cumulative_derivative(): + """The differential function is minus the cumulative derivative.""" + params = (1.0, 100.0, -1.0, -2.0, 10.0) + energies = np.array([2.0, 5.0, 20.0, 50.0]) + step = energies * 1e-5 + + upper = energetics.vector_cum_broken_power_law( + energies + step, *params + ) + lower = energetics.vector_cum_broken_power_law( + energies - step, *params + ) + numerical_density = -(upper - lower) / (2 * step) + density = energetics.vector_diff_broken_power_law(energies, *params) + + np.testing.assert_allclose(density, numerical_density, rtol=1e-8) + + +def test_broken_power_law_array_wrappers_preserve_shape(): + params = (1.0, 100.0, -1.0, -2.0, 10.0) + energies = np.array([[1.0, 5.0], [20.0, 100.0]]) + + cumulative = energetics.array_cum_broken_power_law(energies, *params) + differential = energetics.array_diff_broken_power_law(energies, *params) + + assert cumulative.shape == energies.shape + assert differential.shape == energies.shape + + +def test_broken_power_law_vector_functions_accept_scalar(): + params = (1.0, 100.0, -1.0, -2.0, 10.0) + + cumulative = energetics.vector_cum_broken_power_law(10.0, *params) + differential = energetics.vector_diff_broken_power_law(10.0, *params) + + assert np.isscalar(cumulative) + assert np.isscalar(differential) + + +def test_double_broken_power_law_boundaries_and_monotonicity(): + params = (1.0, 1000.0, -0.5, -1.0, -2.0, 10.0, 100.0) + energies = np.array([0.5, 1.0, 5.0, 10.0, 50.0, 100.0, 500.0, + 1000.0, 2000.0]) + + cumulative = energetics.vector_cum_double_broken_power_law( + energies, *params + ) + + assert cumulative[0] == 1.0 + assert cumulative[1] == 1.0 + assert cumulative[-2] == 0.0 + assert cumulative[-1] == 0.0 + assert np.all(np.diff(cumulative) <= 0) + + +def test_double_broken_power_law_is_continuous_at_breaks(): + params = (1.0, 1000.0, -0.5, -1.0, -2.0, 10.0, 100.0) + for break_energy in (10.0, 100.0): + epsilon = break_energy * 1e-9 + values = energetics.vector_diff_double_broken_power_law( + np.array([break_energy - epsilon, break_energy + epsilon]), + *params, + ) + np.testing.assert_allclose(values[0], values[1], rtol=1e-7) + + +def test_diff_double_broken_power_law_matches_cumulative_derivative(): + params = (1.0, 1000.0, -0.5, -1.0, -2.0, 10.0, 100.0) + energies = np.array([2.0, 5.0, 20.0, 50.0, 200.0, 500.0]) + step = energies * 1e-5 + + upper = energetics.vector_cum_double_broken_power_law( + energies + step, *params + ) + lower = energetics.vector_cum_double_broken_power_law( + energies - step, *params + ) + numerical_density = -(upper - lower) / (2 * step) + density = energetics.vector_diff_double_broken_power_law( + energies, *params + ) + + np.testing.assert_allclose(density, numerical_density, rtol=1e-8) + + +def test_double_broken_power_law_zero_indices(): + params = (1.0, 1000.0, 0.0, 0.0, 0.0, 10.0, 100.0) + cumulative = energetics.vector_cum_double_broken_power_law( + np.array([1.0, 10.0, 100.0, 1000.0]), *params + ) + np.testing.assert_allclose(cumulative, [1.0, 2 / 3, 1 / 3, 0.0]) + + +def test_double_broken_power_law_wrappers_and_scalar(): + params = (1.0, 1000.0, -0.5, -1.0, -2.0, 10.0, 100.0) + energies = np.array([[1.0, 10.0], [100.0, 1000.0]]) + + cumulative = energetics.array_cum_double_broken_power_law( + energies, *params + ) + differential = energetics.array_diff_double_broken_power_law( + energies, *params + ) + assert cumulative.shape == energies.shape + assert differential.shape == energies.shape + assert np.isscalar( + energetics.vector_cum_double_broken_power_law(10.0, *params) + ) + assert np.isscalar( + energetics.vector_diff_double_broken_power_law(10.0, *params) + ) + + +def test_double_broken_power_law_rejects_invalid_energy_order(): + params = (1.0, 1000.0, -0.5, -1.0, -2.0, 100.0, 10.0) + with pytest.raises(ValueError, match="Emin < Eb1 < Eb2 < Emax"): + energetics.vector_cum_double_broken_power_law(10.0, *params) diff --git a/zdm/tests/test_mcmc.py b/zdm/tests/test_mcmc.py index 8742f96d..705e1cb7 100644 --- a/zdm/tests/test_mcmc.py +++ b/zdm/tests/test_mcmc.py @@ -17,6 +17,10 @@ import multiprocessing as mp from multiprocessing import cpu_count +from zdm import MCMC +from zdm import grid +from zdm import parameters + def log_prob(theta): t = time.time() + np.random.uniform(0.005, 0.008) @@ -47,3 +51,127 @@ def test_mcmc(): end = time.time() serial_time = end - start print("MP took {0:.1f} seconds".format(serial_time)) + + +def test_broken_power_law_joint_prior(): + state = parameters.State() + state.energy.luminosity_function = 4 + state.energy.lEmin = 38.0 + state.energy.lEb = 40.0 + state.energy.lEmax = 42.0 + + assert MCMC.valid_parameter_combination({}, state) + assert MCMC.valid_parameter_combination({'lEb': 41.0}, state) + assert not MCMC.valid_parameter_combination({'lEb': 37.0}, state) + assert not MCMC.valid_parameter_combination({'lEmin': 41.0}, state) + assert not MCMC.valid_parameter_combination({'lEmax': 39.0}, state) + + posterior = MCMC.calc_log_posterior( + [37.0], + state, + {'lEb': {'min': 35.0, 'max': 43.0}}, + [[], []], + ) + assert posterior == -np.inf + + +def test_broken_power_law_initial_walkers_obey_joint_prior(): + state = parameters.State() + state.energy.luminosity_function = 4 + params = { + 'lEmin': {'min': 37.0, 'max': 41.0}, + 'lEb': {'min': 38.0, 'max': 42.0}, + 'lEmax': {'min': 39.0, 'max': 43.0}, + 'gamma': {'min': -3.0, 'max': 0.0}, + 'gamma2': {'min': -5.0, 'max': 0.0}, + } + + walkers = MCMC.get_initial_walkers( + state, params, nwalkers=64, rng=np.random.default_rng(1234) + ) + indices = {name: i for i, name in enumerate(params)} + + assert np.all( + walkers[:, indices['lEmin']] + < walkers[:, indices['lEb']] + ) + assert np.all( + walkers[:, indices['lEb']] + < walkers[:, indices['lEmax']] + ) + + +def test_double_broken_power_law_joint_prior(): + state = parameters.State() + state.energy.luminosity_function = 5 + state.energy.lEmin = 37.0 + state.energy.lEb = 39.0 + state.energy.lEb2 = 41.0 + state.energy.lEmax = 43.0 + + assert MCMC.valid_parameter_combination({}, state) + assert MCMC.valid_parameter_combination({'lEb2': 42.0}, state) + assert not MCMC.valid_parameter_combination({'lEb': 42.0}, state) + assert not MCMC.valid_parameter_combination({'lEb2': 38.0}, state) + assert not MCMC.valid_parameter_combination({'lEmin': 40.0}, state) + assert not MCMC.valid_parameter_combination({'lEmax': 40.0}, state) + + +def test_double_broken_power_law_initial_walkers_obey_joint_prior(): + state = parameters.State() + state.energy.luminosity_function = 5 + params = { + 'lEmin': {'min': 36.0, 'max': 40.0}, + 'lEb': {'min': 37.0, 'max': 41.0}, + 'lEb2': {'min': 39.0, 'max': 43.0}, + 'lEmax': {'min': 41.0, 'max': 45.0}, + 'gamma': {'min': -3.0, 'max': 0.0}, + 'gamma2': {'min': -5.0, 'max': 0.0}, + 'gamma3': {'min': -7.0, 'max': 0.0}, + } + + walkers = MCMC.get_initial_walkers( + state, params, nwalkers=64, rng=np.random.default_rng(4321) + ) + indices = {name: i for i, name in enumerate(params)} + + assert np.all( + walkers[:, indices['lEmin']] < walkers[:, indices['lEb']] + ) + assert np.all( + walkers[:, indices['lEb']] < walkers[:, indices['lEb2']] + ) + assert np.all( + walkers[:, indices['lEb2']] < walkers[:, indices['lEmax']] + ) + + +def test_grid_selects_double_broken_power_law_functions(): + state = parameters.State() + state.energy.luminosity_function = 5 + state.energy.lEb = 1.0 + state.energy.lEb2 = 2.0 + state.energy.gamma2 = -1.0 + state.energy.gamma3 = -2.0 + + test_grid = grid.Grid.__new__(grid.Grid) + test_grid.state = state + test_grid.luminosity_function = 5 + test_grid.init_luminosity_functions() + + cumulative = test_grid.vector_cum_lf( + np.array([1.0, 10.0, 100.0, 1000.0]), + 1.0, + 1000.0, + -0.5, + ) + differential = test_grid.vector_diff_lf( + np.array([10.0, 100.0]), + 1.0, + 1000.0, + -0.5, + ) + + assert cumulative[0] == 1.0 + assert cumulative[-1] == 0.0 + assert np.all(differential > 0) diff --git a/zdm/tests/test_parameters.py b/zdm/tests/test_parameters.py index 28b3700d..260bdeee 100644 --- a/zdm/tests/test_parameters.py +++ b/zdm/tests/test_parameters.py @@ -9,4 +9,17 @@ def test_init_state(): # Fuss a bit assert state.analysis.NewGrids -test_init_state() \ No newline at end of file + +def test_broken_power_law_parameters(): + state = parameters.State() + + assert state.energy.gamma2 == -2.0 + assert state.energy.gamma3 == -3.0 + assert state.energy.lEb == 40.0 + assert state.energy.lEb2 == 41.0 + assert state.params["gamma2"] == "energy" + assert state.params["gamma3"] == "energy" + assert state.params["lEb"] == "energy" + assert state.params["lEb2"] == "energy" + +test_init_state() From f33ad2e44b6f20e1bcc711023b0c84861a94a451 Mon Sep 17 00:00:00 2001 From: Clancy James Date: Fri, 28 Aug 2026 10:12:57 +0800 Subject: [PATCH 02/13] resolved fixes from PR --- papers/FitRepetition2025/run_slice_vs_emin.py | 275 ++++++++++++++++++ papers/FitRepetition2025/slurm/run_mcmc.slurm | 49 ++++ zdm/scripts/run_slice.py | 63 +--- zdm/scripts/slurm/run_mcmc.slurm | 4 +- 4 files changed, 338 insertions(+), 53 deletions(-) create mode 100644 papers/FitRepetition2025/run_slice_vs_emin.py create mode 100755 papers/FitRepetition2025/slurm/run_mcmc.slurm diff --git a/papers/FitRepetition2025/run_slice_vs_emin.py b/papers/FitRepetition2025/run_slice_vs_emin.py new file mode 100644 index 00000000..b51cc3f2 --- /dev/null +++ b/papers/FitRepetition2025/run_slice_vs_emin.py @@ -0,0 +1,275 @@ +import argparse +import numpy as np +import os + +from zdm import figures +from zdm import iteration as it + +from zdm import parameters +from zdm import repeat_grid as zdm_repeat_grid +from zdm import MCMC +from zdm import survey +from zdm import misc_functions + +from astropy.cosmology import Planck18 + +import matplotlib.pyplot as plt +import time +from pkg_resources import resource_filename + +#============================================================================== +''' +Function: main +Date: 10/01/2024 +Purpose: + Main function to run the slice calculation +''' +def main(): + + t0 = time.time() + parser = argparse.ArgumentParser() + parser.add_argument(dest='param',type=str,help="Parameter to do the slice in") + parser.add_argument(dest='min',type=float,help="Min value") + parser.add_argument(dest='max',type=float,help="Max value") + parser.add_argument('-f', '--files', default=None, nargs='+', type=str, help="Survey file names") + parser.add_argument('-r', '--rep_surveys', default=None, nargs='+', type=str, help="Surveys to consider repeaters in") + parser.add_argument('-n',dest='n',type=int,default=50,help="Number of values") + # parser.add_argument('-r',dest='repeaters',default=False,action='store_true',help="Surveys are repeater surveys") + args = parser.parse_args() + + # Values to do the slice in + vals = np.linspace(args.min, args.max, args.n) + # vals2 = np.linspace(34, 39, 4) + vals2 = [39.0] + + # Initialisation + state, surveys_sep = init(args) + + # Set the output directory + outdir = 'cube/' + args.param + '/' + if not os.path.exists(outdir): + os.makedirs(outdir) + + for lEmin in vals2: + state.update_param('lEmin', lEmin) + # Do the slice calculation + ll_lists = calc_slice(vals, state, surveys_sep, args) + + # Plot the slice + out = outdir + '/lEmin_2_' + str(lEmin) + '/' + if not os.path.exists(out): + os.makedirs(out) + plot_slice(vals, ll_lists, surveys_sep, args, out) + +#============================================================================== +''' +Function: init +Date: 10/01/2024 +Purpose: + Initialise state and surveys for the slice calculation +''' +def init(args): + # Set state + state = parameters.State() + state.set_astropy_cosmo(Planck18) + # param_dict={'sfr_n': 1.13, 'alpha': 1.5, 'lmean': 2.27, 'lsigma': 0.55, + # 'lEmax': 41.26, 'lEmin': 39.5, 'gamma': -0.95, 'H0': 73, + # 'min_lat': 0.0, 'sigmaDMG': 0.0, 'sigmaHalo': 20.0} + # param_dict={'sfr_n': 0.8808527057055584, 'alpha': 0.7895161131856694, + # 'lmean': 2.1198711983468064, 'lsigma': 0.44944780033763343, + # 'lEmax': 41.18671139482926, 'lEmin': 39.81049090314043, 'gamma': -1.1558450520609953, + # 'H0': 54.6887137195215, 'halo_method': 0, 'sigmaDMG': 0.0, 'sigmaHalo': 0.0, 'min_lat': 30.0} + # param_dict={'sfr_n': 3.1, 'alpha': 1.4859524003747502, + # 'lmean': 2.3007428869522486, 'lsigma': 0.396300210604263, + # 'lEmax': 40.5, 'lEmin': 39, 'gamma': -1.12, + # 'H0': 70.51322705185869, 'DMhalo': 39.800465306883666} + param_dict={'sfr_n': 2.8727580728334483, 'alpha': 1.4311162666594126, + 'lmean': 2.182113926164531, 'lsigma': 0.43672819419999337, + 'lEmax': 40.91165578515364, 'lEmin': 30.0, #38.394926807403984, + 'gamma': -1.1268723802877352, 'H0': 70.6408065355808, 'DMhalo': 61.038340637162705, + 'halo_method': 0, 'sigmaDMG': 0.2, 'sigmaHalo': 15.0, 'min_lat': 20.0} + # param_dict={'lEmax': 40.578551786703116} + + # param_dict={'sfr_n': 2.0968103423638667, + # 'alpha': 1.5849745889763187, + # 'lmean': 2.126075529180481, + # 'lsigma': 0.706259793231814, + # 'lEmin': 38.394926807403984, + # 'lEmax': 40.91165578515364, + # 'gamma': 0.7988426617566275, + # 'H0': 72.01478718920049, + # 'DMhalo': 23.895321681881498, + # 'lRmin': -2.937758379319553, + # 'lRmax': 3.843455475083895, + # 'Rgamma': -2.3067809502252885} + state.update_params(param_dict) + + # state.update_param('Rgamma', -2.2) + # state.update_param('lRmax', 3.0) + # state.update_param('lRmin', -4.0) + # state.update_param('min_lat', 30.0) + + # Initialise surveys + surveys_sep = [[], []] + + zDMgrid, zvals,dmvals = misc_functions.get_zdm_grid( + state, new=True, plot=False, method='analytic', + datdir=resource_filename('zdm', 'GridData')) + + if args.files is not None: + for survey_name in args.files: + s = survey.load_survey(survey_name, state, dmvals, zvals) + surveys_sep[0].append(s) + + if args.rep_surveys is not None: + for survey_name in args.rep_surveys: + s = survey.load_survey(survey_name, state, dmvals, zvals) + surveys_sep[1].append(s) + + # state.update_param('halo_method', 1) + # state.update_param(args.param, vals[0]) + + return state, surveys_sep + +#============================================================================== +''' +Function: calc_slice +Date: 10/01/2024 +Purpose: + Calculate log likelihoods for a slice in parameter space +''' +def calc_slice(vals, state, surveys_sep, args): + ll_lists = [] + for val in vals: + print("val:", val) + param = {args.param: {'min': -np.inf, 'max': np.inf}} + + ll, ll_list = MCMC.calc_log_posterior([val], state, param, surveys_sep, ind_surveys=True, Pn=True, pNreps=True) + print("ll, ll_list:", ll, ll_list, flush=True) + ll_lists.append(ll_list) + print(ll_lists) + ll_lists = np.asarray(ll_lists) + + return ll_lists + +#============================================================================== +''' +Function: plot_slice +Date: 10/01/2024 +Purpose: + Plot the log likelihoods for the slice in parameter space +''' +def plot_slice(vals, ll_lists, surveys_sep, args, outdir): + plt.figure() + plt.clf() + + llsum = np.zeros(ll_lists.shape[0]) + surveys = surveys_sep[0] + surveys_sep[1] + for i in range(len(surveys)): + s = surveys[i] + lls = ll_lists[:, i] + + lls[lls < -1e10] = -np.inf + lls[np.argwhere(np.isnan(lls))] = -np.inf + + llsum += lls + + lls = lls - np.max(lls) + + # plt.figure() + # plt.clf() + plt.plot(vals, lls, label=s.name) + plt.xlabel(args.param) + plt.ylabel('log likelihood') + # plt.savefig(os.path.join(outdir, s.name + ".pdf")) + + print(vals) + print(llsum) + peak=vals[np.argwhere(llsum == np.max(llsum))[0]] + print("peak", peak) + plt.axvline(peak) + plt.legend() + plt.savefig(outdir + args.param + ".pdf") + + # llsum = llsum - np.max(llsum) + # llsum[llsum < -1e10] = -np.inf + plt.figure() + plt.clf() + plt.plot(vals, llsum, label='Total') + plt.axvline(peak) + # plt.plot(vals, llsum2) + plt.xlabel(args.param) + plt.ylabel('log likelihood') + plt.legend() + + plt.savefig(outdir + args.param + "_sum.pdf") + + np.save(outdir + args.param + "_vals.npy", vals) + np.save(outdir + args.param + "_lls2.npy", llsum) + +#============================================================================== +""" +Function: plot_grids +Date: 10/01/2024 +Purpose: + Plot grids. Adapted from zdm/scripts/plot_pzdm_grid.py + +Imports: + grids = list of grids + surveys = list of surveys + outdir = output directory + val = parameter value for this grid +""" +def plot_grids(grids, surveys, outdir, val): + for g,s in zip(grids, surveys): + zvals=[] + dmvals=[] + nozlist=[] + + if s.zlist is not None: + for iFRB in s.zlist: + zvals.append(s.Zs[iFRB]) + dmvals.append(s.DMEGs[iFRB]) + if s.nozlist is not None: + for dm in s.DMEGs[s.nozlist]: + nozlist.append(dm) + + frbzvals = np.array(zvals) + frbdmvals = np.array(dmvals) + + figures.plot_grid( + g.rates, + g.zvals, + g.dmvals, + name=outdir + s.name + "_" + str(val) + ".pdf", + norm=3, + log=True, + label="$\\log_{10} p({\\rm DM}_{\\rm EG},z)$ [a.u.]", + project=False, + FRBDM=frbdmvals, + FRBZ=frbzvals, + Aconts=[0.01, 0.1, 0.5], + zmax=1.5, + DMmax=3000, + # DMlines=nozlist, + ) + +#============================================================================== +""" +Function: commasep +Date: 23/08/2022 +Purpose: + Turn a string of variables seperated by commas into a list + +Imports: + s = String of variables + +Exports: + List conversion of s +""" +def commasep(s): + return list(map(str, s.split(','))) + +#============================================================================== + +main() diff --git a/papers/FitRepetition2025/slurm/run_mcmc.slurm b/papers/FitRepetition2025/slurm/run_mcmc.slurm new file mode 100755 index 00000000..16986914 --- /dev/null +++ b/papers/FitRepetition2025/slurm/run_mcmc.slurm @@ -0,0 +1,49 @@ +#!/bin/bash +#SBATCH --job-name=reps_Pn_new_3 +#SBATCH --output=../../mcmc/reps_Pn_new_3.out +#SBATCH --ntasks=1 +#SBATCH --cpus-per-task=20 +#SBATCH --time=48:00:00 +#SBATCH --export=NONE +#SBATCH --mem=32GB +# SBATCH --mem-per-cpu=8GB + +############################################################################### +# Author: Jordan Hoffmann # +# Date: 04/06/2024 # +# Purpose: # +# Slurm script for an MCMC run. # +# Usage: # +# Change job-name and output in SBATCH commands # +# Change outfile # +# Check surveys to be used (assumed to be in default survey location) # +# Check command line parameters to run MCMC_wrap2.py # +############################################################################### + +source $ZDM/.venv/bin/activate + +cd $ZDM/zdm + +outfile="mcmc/reps_Pn_new" +walkers=40 +steps=3000 + +# surveys="DSA no_Tobs/MeerTRAPcoherent no_Tobs/MeerTRAPincoherent no_Tobs/FAST no_Tobs/CRAFT_class_I_and_II no_Tobs/parkes_mb_class_I_and_II" +surveys="DSA_34 MeerTRAPcoherent MeerTRAPincoherent FAST CRAFT_class_I_and_II parkes_mb_class_I_and_II" + +rep_surveys="CRAFT_average_ICS CHIME/CHIME_decbin_0_of_6 CHIME/CHIME_decbin_1_of_6 CHIME/CHIME_decbin_2_of_6 CHIME/CHIME_decbin_3_of_6 CHIME/CHIME_decbin_4_of_6 CHIME/CHIME_decbin_5_of_6" +# rep_surveys=CHIME/CHIME_decbin_3_of_6 +# cd data/Surveys/ +# rep_surveys=$(ls CHIME/*) +# rep_surveys=${rep_surveys//".ecsv"/""} +# cd $ZDM/zdm + +echo "Outfile: $outfile.h5" +echo "Walkers: $walkers" +echo "Steps: $steps" + +# command="python scripts/MCMC/MCMC_wrap.py -f $surveys -p data/MCMC/params.json -o $outfile -w $walkers -s $steps" +command="srun python scripts/MCMC/MCMC_wrap.py -f $surveys -r $rep_surveys -p data/MCMC/params2.json -o $outfile -w $walkers -s $steps --Pn --pwb" +# command="python scripts/MCMC/MCMC_wrap.py -f $surveys -r $rep_surveys -p data/MCMC/params3.json -o $outfile -w $walkers -s $steps --Pn" +echo $command +$command \ No newline at end of file diff --git a/zdm/scripts/run_slice.py b/zdm/scripts/run_slice.py index 245c84cc..292e8914 100644 --- a/zdm/scripts/run_slice.py +++ b/zdm/scripts/run_slice.py @@ -23,6 +23,7 @@ Date: 10/01/2024 Purpose: Main function to run the slice calculation + Example: python3 run_slice.py H0 50 80 -n 7 -f CRAFT_ICS_1300 ''' def main(): @@ -39,27 +40,20 @@ def main(): # Values to do the slice in vals = np.linspace(args.min, args.max, args.n) - # vals2 = np.linspace(34, 39, 4) - vals2 = [39.0] + # Initialisation state, surveys_sep = init(args) # Set the output directory - outdir = 'cube/' + args.param + '/' + outdir = 'Slices/' + args.param + '/' if not os.path.exists(outdir): os.makedirs(outdir) - - for lEmin in vals2: - state.update_param('lEmin', lEmin) - # Do the slice calculation - ll_lists = calc_slice(vals, state, surveys_sep, args) - - # Plot the slice - out = outdir + '/lEmin_2_' + str(lEmin) + '/' - if not os.path.exists(out): - os.makedirs(out) - plot_slice(vals, ll_lists, surveys_sep, args, out) + + # Do the slice calculation + ll_lists = calc_slice(vals, state, surveys_sep, args) + + plot_slice(vals, ll_lists, surveys_sep, args, outdir) #============================================================================== ''' @@ -72,43 +66,14 @@ def init(args): # Set state state = parameters.State() state.set_astropy_cosmo(Planck18) - # param_dict={'sfr_n': 1.13, 'alpha': 1.5, 'lmean': 2.27, 'lsigma': 0.55, - # 'lEmax': 41.26, 'lEmin': 39.5, 'gamma': -0.95, 'H0': 73, - # 'min_lat': 0.0, 'sigmaDMG': 0.0, 'sigmaHalo': 20.0} - # param_dict={'sfr_n': 0.8808527057055584, 'alpha': 0.7895161131856694, - # 'lmean': 2.1198711983468064, 'lsigma': 0.44944780033763343, - # 'lEmax': 41.18671139482926, 'lEmin': 39.81049090314043, 'gamma': -1.1558450520609953, - # 'H0': 54.6887137195215, 'halo_method': 0, 'sigmaDMG': 0.0, 'sigmaHalo': 0.0, 'min_lat': 30.0} - # param_dict={'sfr_n': 3.1, 'alpha': 1.4859524003747502, - # 'lmean': 2.3007428869522486, 'lsigma': 0.396300210604263, - # 'lEmax': 40.5, 'lEmin': 39, 'gamma': -1.12, - # 'H0': 70.51322705185869, 'DMhalo': 39.800465306883666} param_dict={'sfr_n': 2.8727580728334483, 'alpha': 1.4311162666594126, 'lmean': 2.182113926164531, 'lsigma': 0.43672819419999337, 'lEmax': 40.91165578515364, 'lEmin': 30.0, #38.394926807403984, 'gamma': -1.1268723802877352, 'H0': 70.6408065355808, 'DMhalo': 61.038340637162705, 'halo_method': 0, 'sigmaDMG': 0.2, 'sigmaHalo': 15.0, 'min_lat': 20.0} - # param_dict={'lEmax': 40.578551786703116} - - # param_dict={'sfr_n': 2.0968103423638667, - # 'alpha': 1.5849745889763187, - # 'lmean': 2.126075529180481, - # 'lsigma': 0.706259793231814, - # 'lEmin': 38.394926807403984, - # 'lEmax': 40.91165578515364, - # 'gamma': 0.7988426617566275, - # 'H0': 72.01478718920049, - # 'DMhalo': 23.895321681881498, - # 'lRmin': -2.937758379319553, - # 'lRmax': 3.843455475083895, - # 'Rgamma': -2.3067809502252885} + state.update_params(param_dict) - - # state.update_param('Rgamma', -2.2) - # state.update_param('lRmax', 3.0) - # state.update_param('lRmin', -4.0) - # state.update_param('min_lat', 30.0) - + # Initialise surveys surveys_sep = [[], []] @@ -125,10 +90,7 @@ def init(args): for survey_name in args.rep_surveys: s = survey.load_survey(survey_name, state, dmvals, zvals) surveys_sep[1].append(s) - - # state.update_param('halo_method', 1) - # state.update_param(args.param, vals[0]) - + return state, surveys_sep #============================================================================== @@ -250,8 +212,7 @@ def plot_grids(grids, surveys, outdir, val): FRBZ=frbzvals, Aconts=[0.01, 0.1, 0.5], zmax=1.5, - DMmax=3000, - # DMlines=nozlist, + DMmax=3000 ) #============================================================================== diff --git a/zdm/scripts/slurm/run_mcmc.slurm b/zdm/scripts/slurm/run_mcmc.slurm index 16986914..24f02bf0 100755 --- a/zdm/scripts/slurm/run_mcmc.slurm +++ b/zdm/scripts/slurm/run_mcmc.slurm @@ -29,7 +29,7 @@ walkers=40 steps=3000 # surveys="DSA no_Tobs/MeerTRAPcoherent no_Tobs/MeerTRAPincoherent no_Tobs/FAST no_Tobs/CRAFT_class_I_and_II no_Tobs/parkes_mb_class_I_and_II" -surveys="DSA_34 MeerTRAPcoherent MeerTRAPincoherent FAST CRAFT_class_I_and_II parkes_mb_class_I_and_II" +surveys="DSA MeerTRAPcoherent MeerTRAPincoherent FAST CRAFT_class_I_and_II parkes_mb_class_I_and_II" rep_surveys="CRAFT_average_ICS CHIME/CHIME_decbin_0_of_6 CHIME/CHIME_decbin_1_of_6 CHIME/CHIME_decbin_2_of_6 CHIME/CHIME_decbin_3_of_6 CHIME/CHIME_decbin_4_of_6 CHIME/CHIME_decbin_5_of_6" # rep_surveys=CHIME/CHIME_decbin_3_of_6 @@ -46,4 +46,4 @@ echo "Steps: $steps" command="srun python scripts/MCMC/MCMC_wrap.py -f $surveys -r $rep_surveys -p data/MCMC/params2.json -o $outfile -w $walkers -s $steps --Pn --pwb" # command="python scripts/MCMC/MCMC_wrap.py -f $surveys -r $rep_surveys -p data/MCMC/params3.json -o $outfile -w $walkers -s $steps --Pn" echo $command -$command \ No newline at end of file +$command From a57d27315ae3bfe91bb1bca37024b1cf97eb3545 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=90=B4=E6=B2=81?= Date: Fri, 28 Aug 2026 15:48:33 +0800 Subject: [PATCH 03/13] Add Break-Schechter luminosity function --- zdm/MCMC.py | 4 +- zdm/data/MCMC/break_schechter.json | 92 +++++++++++++++ zdm/energetics.py | 180 ++++++++++++++++++++++++++++- zdm/grid.py | 32 ++++- zdm/iteration.py | 9 ++ zdm/parameters.py | 6 +- zdm/tests/test_energetics.py | 48 ++++++++ zdm/tests/test_mcmc.py | 26 +++++ 8 files changed, 388 insertions(+), 9 deletions(-) create mode 100644 zdm/data/MCMC/break_schechter.json diff --git a/zdm/MCMC.py b/zdm/MCMC.py index eb5b18d9..606ad6fd 100644 --- a/zdm/MCMC.py +++ b/zdm/MCMC.py @@ -63,7 +63,7 @@ def valid_parameter_combination(param_dict, state): ) lEmin = param_dict.get('lEmin', state.energy.lEmin) lEmax = param_dict.get('lEmax', state.energy.lEmax) - if luminosity_function == 4: + if luminosity_function in (4, 6): lEb = param_dict.get('lEb', state.energy.lEb) return bool(lEmin < lEb < lEmax) if luminosity_function == 5: @@ -95,7 +95,7 @@ def get_initial_walkers(state, params, nwalkers, rng=None, max_attempts=10000): else: raise ValueError( "Could not initialize MCMC walkers inside the joint priors. " - "For luminosity_function=4 or 5, ensure the prior ranges " + "For luminosity_function=4, 5, or 6, ensure the prior ranges " "permit the required ordering of break energies." ) diff --git a/zdm/data/MCMC/break_schechter.json b/zdm/data/MCMC/break_schechter.json new file mode 100644 index 00000000..d7da562f --- /dev/null +++ b/zdm/data/MCMC/break_schechter.json @@ -0,0 +1,92 @@ +{ + "mcmc": { + "parameter_order": [ + "sfr_n", + "alpha", + "lmean", + "lsigma", + "lEmin", + "lEb", + "lEmax", + "gamma", + "gamma2", + "H0" + ] + }, + "config": { + "luminosity_function": 6, + "alpha_method": 1, + "source_evolution": 0, + "logF": -0.494850021680094, + "DMhalo": 50.0, + "sigmaDMG": 0.2, + "sigmaHalo": 15.0, + "halo_method": 0, + "Wlogmean": 0.0, + "Wlogsigma": 0.42, + "WNbins": 5, + "WidthFunction": 1, + "Wthresh": 0.5, + "Wmethod": 2, + "WMin": 0.1, + "WMax": 100.0, + "WNInternalBins": 100, + "Slogmean": 0.305, + "Slogsigma": 0.75, + "ScatFunction": 1, + "Sbackproject": true, + "Sfnorm": 600.0, + "Sfpower": -4.0, + "Smaxsigma": 3.0 + }, + "sfr_n": { + "DC": "FRBdemo", + "min": -2.0, + "max": 4.0 + }, + "alpha": { + "DC": "energy", + "min": -0.5, + "max": 2.5 + }, + "lmean": { + "DC": "host", + "min": 1.0, + "max": 3.0 + }, + "lsigma": { + "DC": "host", + "min": 0.1, + "max": 1.5 + }, + "lEmin": { + "DC": "energy", + "min": 36.0, + "max": 40.5 + }, + "lEb": { + "DC": "energy", + "min": 36.0, + "max": 45.0 + }, + "lEmax": { + "DC": "energy", + "min": 40.5, + "max": 45.0 + }, + "gamma": { + "DC": "energy", + "min": -3.0, + "max": 0.0 + }, + "gamma2": { + "DC": "energy", + "min": -3.0, + "max": 0.0 + }, + "H0": { + "DC": "cosmo", + "min": 35.0, + "max": 110.0 + } +} \ No newline at end of file diff --git a/zdm/energetics.py b/zdm/energetics.py index f3fce691..98c539cb 100644 --- a/zdm/energetics.py +++ b/zdm/energetics.py @@ -7,6 +7,10 @@ 1. **Power Law**: Simple power-law distribution dN/dE ~ E^gamma between Emin and Emax 2. **Gamma Function**: Upper incomplete gamma function distribution with exponential cutoff +3. **Break Power Law**: A continuous two-segment power law with one break energy Eb +4. **Two Break Power Law**: A continuous three-segment power law with two break energies +5. **Break-Schechter Function**: A power law with slope gamma1 below the break energy Eb + and a Schechter-like distribution with slope gamma2 and exponential cutoff above Eb The gamma function implementation uses spline interpolation for efficiency, as direct evaluation of the incomplete gamma function is computationally expensive @@ -17,6 +21,7 @@ - `vector_cum_power_law`: Cumulative power-law luminosity function - `vector_cum_broken_power_law`: Cumulative broken power-law luminosity function - `vector_cum_double_broken_power_law`: Cumulative two-break power-law luminosity function +- `vector_cum_broken_gamma`: Cumulative Break-Schechter luminosity function. - `vector_cum_gamma_spline`: Cumulative gamma function with spline interpolation - `array_cum_gamma_spline`: N-dimensional array wrapper for gamma function - `init_igamma_splines`: Initialize spline lookup tables for gamma functions @@ -503,6 +508,178 @@ def array_diff_double_broken_power_law(Eth, *params): return result.reshape(dims) +########### broken Schechter functions ############# + +def _broken_schechter_upper_gamma(x, gamma): + """ + Fast upper incomplete Gamma(gamma, x), using the existing spline cache. + """ + global SplineLog + + x = np.asarray(x, dtype=float) + scalar_input = x.ndim == 0 + x = np.atleast_1d(x) + + if np.any(x <= 0.0): + raise ValueError("Upper incomplete gamma arguments must be positive") + + if gamma not in igamma_splines: + init_igamma_splines([gamma]) + + if SplineLog: + result = 10 ** interpolate.splev(np.log10(x), igamma_splines[gamma]) + else: + result = interpolate.splev(x, igamma_splines[gamma]) + + result = np.asarray(result, dtype=float) + + if scalar_input: + return result[0] + + return result + + +def _broken_schechter_lower_integral(x, gamma): + """ + Integral of u^(gamma - 1) from x to 1. + + Equal to (1 - x**gamma)/gamma, with the gamma=0 logarithmic limit. + """ + x = np.asarray(x, dtype=float) + + if np.isclose(gamma, 0.0): + return -np.log(x) + + return -np.expm1(gamma * np.log(x)) / gamma + + +def _broken_schechter_normalization(Emin, Ecut, gamma1, gamma2, Eb): + """ + Dimensionless normalization for the broken-Schechter distribution. + """ + low_mass = _broken_schechter_lower_integral(Emin / Eb, gamma1) + + high_scale = (np.exp(Eb / Ecut) * (Ecut / Eb) ** gamma2) + high_mass = (high_scale * _broken_schechter_upper_gamma(Eb / Ecut, gamma2)) + + return low_mass + high_mass, high_scale, high_mass + + +def vector_cum_broken_schechter(Eth, *params): + """ + Survival function P(E > Eth) for a broken-Schechter energy function. + + Parameters + ---------- + Eth : scalar or one-dimensional ndarray + Energy threshold in erg. + params : tuple + (Emin, Ecut, gamma1, gamma2, Eb) + + Notes + ----- + Ecut corresponds to the existing zdm Gamma parameter called Emax. + It is a characteristic exponential cutoff energy, not a hard maximum. + """ + Emin, Ecut, gamma1, gamma2, Eb = params + + if not (0.0 < Emin < Eb < Ecut): + raise ValueError("Broken Schechter energies must satisfy 0 < Emin < Eb < Ecut") + + Eth = np.asarray(Eth, dtype=float) + scalar_input = Eth.ndim == 0 + + if Eth.ndim > 1: + raise ValueError("Eth must be scalar or one-dimensional") + + Eth = np.atleast_1d(Eth) + + normalization, high_scale, high_at_break = (_broken_schechter_normalization(Emin, Ecut, gamma1, gamma2, Eb)) + + result = np.empty_like(Eth) + + below_min = Eth < Emin + below_break = (Eth >= Emin) & (Eth < Eb) + above_break = Eth >= Eb + + result[below_min] = 1.0 + result[below_break] = (_broken_schechter_lower_integral(Eth[below_break] / Eb, gamma1) + high_at_break) / normalization + + if np.any(above_break): + result[above_break] = (high_scale * _broken_schechter_upper_gamma(Eth[above_break] / Ecut, gamma2) / normalization) + + result = np.clip(result, 0.0, 1.0) + + if scalar_input: + return result[0] + + return result + + +def array_cum_broken_schechter(Eth, *params): + """ + N-dimensional wrapper for vector_cum_broken_schechter. + """ + Eth = np.asarray(Eth, dtype=float) + original_shape = Eth.shape + + result = vector_cum_broken_schechter(Eth.ravel(), *params,) + + return result.reshape(original_shape) + + +def vector_diff_broken_schechter(E, *params): + """ + Normalized differential broken-Schechter distribution dP/dE. + + Parameters + ---------- + E : scalar or one-dimensional ndarray + Energy in erg. + params : tuple + (Emin, Ecut, gamma1, gamma2, Eb) + """ + Emin, Ecut, gamma1, gamma2, Eb = params + + if not (0.0 < Emin < Eb < Ecut): + raise ValueError("Broken Schechter energies must satisfy 0 < Emin < Eb < Ecut") + + E = np.asarray(E, dtype=float) + scalar_input = E.ndim == 0 + + if E.ndim > 1: + raise ValueError("E must be scalar or one-dimensional") + + E = np.atleast_1d(E) + + normalization, _, _ = _broken_schechter_normalization(Emin, Ecut, gamma1, gamma2, Eb) + + result = np.zeros_like(E) + + below_break = (E >= Emin) & (E < Eb) + above_break = E >= Eb + + result[below_break] = ((E[below_break] / Eb) ** (gamma1 - 1.0) / (Eb * normalization)) + result[above_break] = ((E[above_break] / Eb) ** (gamma2 - 1.0) * np.exp(-(E[above_break] - Eb) / Ecut) / (Eb * normalization)) + + if scalar_input: + return result[0] + + return result + + +def array_diff_broken_schechter(E, *params): + """ + N-dimensional wrapper for vector_diff_broken_schechter. + """ + E = np.asarray(E, dtype=float) + original_shape = E.shape + + result = vector_diff_broken_schechter(E.ravel(), *params,) + + return result.reshape(original_shape) + + ########### gamma functions ############# @@ -538,8 +715,7 @@ def vector_cum_gamma(Eth, *params): norm = float(mpmath.gammainc(gamma, a=Emin/Emax)) Eth_Emax = Eth/Emax # If this is too slow, we can adopt scipy + recurrance - numer = np.array([float(mpmath.gammainc( - gamma, a=iEE)) for iEE in Eth_Emax]) + numer = np.array([float(mpmath.gammainc(gamma, a=iEE)) for iEE in Eth_Emax]) result=numer/norm # Low end diff --git a/zdm/grid.py b/zdm/grid.py index e41847dc..83544aba 100644 --- a/zdm/grid.py +++ b/zdm/grid.py @@ -178,9 +178,14 @@ def init_luminosity_functions(self): self.vector_cum_lf = self._vector_cum_double_broken_power_law self.array_diff_lf = self._array_diff_double_broken_power_law self.vector_diff_lf = self._vector_diff_double_broken_power_law + elif self.luminosity_function == 6: #Broken Schechter + self.array_cum_lf = self._array_cum_broken_schechter + self.vector_cum_lf = self._vector_cum_broken_schechter + self.array_diff_lf = self._array_diff_broken_schechter + self.vector_diff_lf = self._vector_diff_broken_schechter else: raise ValueError( - "Luminosity function must be one of 0, 1, 2, 3, 4, or 5; " + "Luminosity function must be one of 0, 1, 2, 3, 4, 5, or 6; " f"got {self.luminosity_function}" ) @@ -246,6 +251,29 @@ def _vector_diff_double_broken_power_law( params = self._double_broken_power_law_params(Emin, Emax, gamma1) return energetics.vector_diff_double_broken_power_law(Eth, *params) + def _broken_schechter_params(self, Emin, Ecut, gamma1): + """ + Expand standard grid LF arguments for a broken-Schechter function. + """ + return (Emin,Ecut,gamma1,self.state.energy.gamma2,10 ** self.state.energy.lEb) + + def _array_cum_broken_schechter(self,Eth,Emin,Ecut,gamma1,*unused,): + params = self._broken_schechter_params(Emin, Ecut, gamma1) + return energetics.array_cum_broken_schechter(Eth,*params) + + def _vector_cum_broken_schechter(self,Eth,Emin,Ecut,gamma1,*unused): + params = self._broken_schechter_params(Emin,Ecut,gamma1) + return energetics.vector_cum_broken_schechter(Eth, *params) + + def _array_diff_broken_schechter(self,E,Emin,Ecut,gamma1,*unused): + params = self._broken_schechter_params(Emin,Ecut,gamma1) + return energetics.array_diff_broken_schechter(E,*params) + + def _vector_diff_broken_schechter(self,E,Emin,Ecut,gamma1,*unused): + params = self._broken_schechter_params(Emin,Ecut,gamma1) + return energetics.vector_diff_broken_schechter(E,*params) + + def parse_grid(self, zDMgrid, zvals, dmvals): self.grid = zDMgrid self.zvals = zvals @@ -789,7 +817,7 @@ def GenMCSample(self, N, Poisson=False): """ # Boost? - if self.state.energy.luminosity_function in [1, 2]: + if self.state.energy.luminosity_function in [1, 2, 6]: Emax_boost = 3.0 else: Emax_boost = 0.0 diff --git a/zdm/iteration.py b/zdm/iteration.py index f242588d..8f68cdc1 100644 --- a/zdm/iteration.py +++ b/zdm/iteration.py @@ -1538,6 +1538,15 @@ def ConvertToMeaningfulConstant(state,Eref=1e39): 10 ** state.energy.lEb, 10 ** state.energy.lEb2, ) + elif state.energy.luminosity_function == 6: + factor = energetics.vector_cum_broken_schechter( + np.array([Eref]), + Emin, + Emax, + gamma, + state.energy.gamma2, + 10 ** state.energy.lEb, + ) else: factor = energetics.vector_cum_gamma(np.array([Eref]),Emin,Emax,gamma) const *= factor diff --git a/zdm/parameters.py b/zdm/parameters.py index ba01c457..5b1144ca 100644 --- a/zdm/parameters.py +++ b/zdm/parameters.py @@ -443,7 +443,7 @@ class EnergeticsParams(data_class.myDataClass): gamma2: float = field( default=-2.0, metadata={ - "help": "second slope of a broken power-law luminosity function", + "help": "second slope of a broken power law, or Schechter-segment shape index for a broken-Schechter function", "unit": "", "Notation": "\\gamma_2", }, @@ -459,7 +459,7 @@ class EnergeticsParams(data_class.myDataClass): lEb: float = field( default=40.0, metadata={ - "help": "$\\log_{10}$ of the break energy for a broken power law", + "help": "$\\log_{10}$ of the break energy for a broken power law or broken-Schechter function", "unit": "erg", "Notation": "\\log_{10} E_{\\rm b}", }, @@ -475,7 +475,7 @@ class EnergeticsParams(data_class.myDataClass): luminosity_function: int = field( default=2, metadata={ - "help": "luminosity function applied (0=power-law, 1=gamma, 2=spline+gamma, 3=gamma+linear+log10, 4=broken power-law, 5=doubly-broken power-law)" + "help": "luminosity function applied (0=power-law, 1=gamma, 2=spline+gamma, 3=gamma+linear+log10, 4=broken power-law, 5=doubly-broken power-law, 6=broken-Schechter)", }, ) diff --git a/zdm/tests/test_energetics.py b/zdm/tests/test_energetics.py index 853326c5..1d087def 100644 --- a/zdm/tests/test_energetics.py +++ b/zdm/tests/test_energetics.py @@ -187,3 +187,51 @@ def test_double_broken_power_law_rejects_invalid_energy_order(): params = (1.0, 1000.0, -0.5, -1.0, -2.0, 100.0, 10.0) with pytest.raises(ValueError, match="Emin < Eb1 < Eb2 < Emax"): energetics.vector_cum_double_broken_power_law(10.0, *params) + +def test_broken_schechter_boundaries_and_monotonicity(): + params = (1e37,1e42,-0.5,-1.2,1e40) + energies = np.logspace(36, 45, 300) + + cumulative = energetics.vector_cum_broken_schechter(energies, *params) + + assert cumulative[0] == 1.0 + assert np.all(cumulative >= 0.0) + assert np.all(cumulative <= 1.0) + assert np.all(np.diff(cumulative) <= 1e-12) + assert cumulative[-1] < 1e-10 + + +def test_broken_schechter_is_continuous_at_break(): + params = (1e37,1e42,-0.5,-1.2,1e40) + + Eb = params[-1] + epsilon = Eb * 1e-8 + + values = energetics.vector_diff_broken_schechter(np.array([Eb - epsilon, Eb + epsilon]), *params) + np.testing.assert_allclose(values[0], values[1], rtol=1e-6) + + +def test_broken_schechter_diff_matches_cumulative_derivative(): + params = (1e37,1e42,-0.5,-1.2,1e40) + + energies = np.array([2e37,1e39,2e40,1e42]) + step = energies * 1e-5 + + upper = energetics.vector_cum_broken_schechter(energies + step, *params) + lower = energetics.vector_cum_broken_schechter(energies - step, *params) + + numerical_density = -(upper - lower) / (2.0 * step) + density = energetics.vector_diff_broken_schechter(energies, *params) + + np.testing.assert_allclose(density, numerical_density, rtol=1e-5) + + +def test_broken_schechter_array_wrappers(): + params = (1e37,1e42,-0.5,-1.2,1e40) + + energies = np.array([[1e37, 1e39], [1e40, 1e42]]) + cumulative = energetics.array_cum_broken_schechter(energies, *params) + differential = energetics.array_diff_broken_schechter(energies, *params) + + assert cumulative.shape == energies.shape + assert differential.shape == energies.shape \ No newline at end of file diff --git a/zdm/tests/test_mcmc.py b/zdm/tests/test_mcmc.py index 705e1cb7..26561b84 100644 --- a/zdm/tests/test_mcmc.py +++ b/zdm/tests/test_mcmc.py @@ -175,3 +175,29 @@ def test_grid_selects_double_broken_power_law_functions(): assert cumulative[0] == 1.0 assert cumulative[-1] == 0.0 assert np.all(differential > 0) + + +def test_broken_schechter_joint_prior(): + state = parameters.State() + + state.energy.luminosity_function = 6 + state.energy.lEmin = 38.0 + state.energy.lEb = 40.0 + state.energy.lEmax = 42.0 + + assert MCMC.valid_parameter_combination({}, state) + + assert MCMC.valid_parameter_combination( + {"lEb": 41.0}, + state, + ) + + assert not MCMC.valid_parameter_combination( + {"lEb": 37.0}, + state, + ) + + assert not MCMC.valid_parameter_combination( + {"lEb": 43.0}, + state, + ) \ No newline at end of file From 88aef748ae673f38f342408050a81777e40d9cb6 Mon Sep 17 00:00:00 2001 From: Clancy James Date: Sat, 29 Aug 2026 06:20:36 +0800 Subject: [PATCH 04/13] fixed bad character in Meertrap coherent file, and unwound profiling in MCMC --- papers/FitRepetition2025/profiled.py | 31 +++++++++++++++++++++ zdm/MCMC.py | 37 ++++---------------------- zdm/data/Surveys/MeerTRAPcoherent.ecsv | 4 +-- 3 files changed, 38 insertions(+), 34 deletions(-) create mode 100644 papers/FitRepetition2025/profiled.py diff --git a/papers/FitRepetition2025/profiled.py b/papers/FitRepetition2025/profiled.py new file mode 100644 index 00000000..b787a8e8 --- /dev/null +++ b/papers/FitRepetition2025/profiled.py @@ -0,0 +1,31 @@ +PROFILED_PID = None +import cProfile +from zdm import MCMC + +def profiled_calc_log_posterior(param_vals, state, params, surveys_sep, Pn=False, Pns=False, Pnr=False, + pNreps=True, psnr=True, ptauw=False, pwb=False, + log_halo=False, lin_host=False, ind_surveys=False, g0info=None, nz=500, ndm=1400, + zmax=5.,dmmax=7000., + dopath=False, opstate=None, opt_params=None, opt_model=None): + + global PROFILED_PID + pid = os.getpid() + + # If we haven't chosen a worker yet, choose this one + if PROFILED_PID is None: + PROFILED_PID = pid + + if pid == PROFILED_PID: + profiler_output = f"worker_{pid}.prof" + return cProfile.runctx( + "calc_log_posterior(param_vals, state, params, surveys_sep, Pn, Pns, Pnr, " + "pNreps, psnr, ptauw, pwb, log_halo, lin_host, ind_surveys, g0info, nz, ndm, zmax,dmmax, " + "dopath, opstate, opt_params, opt_model)", + globals(), + locals(), + profiler_output + ) + else: + return MCMC.calc_log_posterior(param_vals, state, params, surveys_sep, Pn, Pns, Pnr, + pNreps, psnr, ptauw, pwb, log_halo, lin_host, ind_surveys, g0info, nz, ndm, zmax,dmmax, + dopath, opstate, opt_params, opt_model) diff --git a/zdm/MCMC.py b/zdm/MCMC.py index 0d54bc63..80c0ed19 100644 --- a/zdm/MCMC.py +++ b/zdm/MCMC.py @@ -51,42 +51,12 @@ from zdm import misc_functions as mf from zdm import repeat_grid import os -import cProfile from zdm import optical_numerics as on from zdm import optical as opt from zdm import optical_params as op #============================================================================== -PROFILED_PID = None - -def profiled_calc_log_posterior(param_vals, state, params, surveys_sep, Pn=False, Pns=False, Pnr=False, - pNreps=True, psnr=True, ptauw=False, pwb=False, - log_halo=False, lin_host=False, ind_surveys=False, g0info=None, nz=500, ndm=1400, - zmax=5.,dmmax=7000., - dopath=False, opstate=None, opt_params=None, opt_model=None): - - global PROFILED_PID - pid = os.getpid() - - # If we haven't chosen a worker yet, choose this one - if PROFILED_PID is None: - PROFILED_PID = pid - - if pid == PROFILED_PID: - profiler_output = f"worker_{pid}.prof" - return cProfile.runctx( - "calc_log_posterior(param_vals, state, params, surveys_sep, Pn, Pns, Pnr, " - "pNreps, psnr, ptauw, pwb, log_halo, lin_host, ind_surveys, g0info, nz, ndm, zmax,dmmax, " - "dopath, opstate, opt_params, opt_model)", - globals(), - locals(), - profiler_output - ) - else: - return calc_log_posterior(param_vals, state, params, surveys_sep, Pn, Pns, Pnr, - pNreps, psnr, ptauw, pwb, log_halo, lin_host, ind_surveys, g0info, nz, ndm, zmax,dmmax, - dopath, opstate, opt_params, opt_model) def calc_log_posterior(param_vals, state, params, surveys_sep, Pn=False, Pns=False, Pnr=False, pNreps=True, psnr=True, ptauw=False, pwb=False, @@ -416,8 +386,11 @@ def mcmc_runner(logpf, outfile, state, params, surveys, nwalkers=10, nsteps=100, # may or may not be needed #os.environ["OMP_NUM_THREADS"] = "1" - cpus = int(os.environ.get("SLURM_CPUS_PER_TASK", 1)) - print(f"Using {cpus} CPUs from Slurm allocation") + if "SLURM_CPUS_PER_TASK" in os.environ: + cpus = int(os.environ.get("SLURM_CPUS_PER_TASK", 1)) + print(f"Using {cpus} CPUs from Slurm allocation") + else: + cpus = None Pool = mp.get_context('fork').Pool # num_cpus = mp.cpu_count() diff --git a/zdm/data/Surveys/MeerTRAPcoherent.ecsv b/zdm/data/Surveys/MeerTRAPcoherent.ecsv index 7fd09f92..705a5ffe 100644 --- a/zdm/data/Surveys/MeerTRAPcoherent.ecsv +++ b/zdm/data/Surveys/MeerTRAPcoherent.ecsv @@ -30,9 +30,9 @@ TNS BW DM DMG FBAR FRES Gb Gl SNR SNRTHRESH 20230306F 770 689.5 23 1284 0.836 76.91 287.65 11.02 8.0 0.066 0.30624 3.6 +14:54:17 12:24:01 -1.0 20230413C 544 1532.2 45 816 0.836 -30.92 252.63 14.8 8.0 0.066 0.30624 19.28 -39:50:00 05:13:00 -1.0 20230613A 770 483.51 30 1284 0.836 -71.61 353.52 49.88 8.0 0.066 0.30624 1.84 -27:03:10.01 23:47:24.65 0.3923 -20230808F 856 653.2 36 1284 0.836 -50.95 263.52 35.78 8.0 0.066 0.30624 7.941 −51:56:07.02 03:33:12.99 0.3472 +20230808F 856 653.2 36 1284 0.836 -50.95 263.52 35.78 8.0 0.066 0.30624 7.941 -51:56:07.02 03:33:12.99 0.3472 20230827E 544 1433.7 38 816 0.836 -42.84 222.57 54.8 8.0 0.066 0.30624 9.6 -18:16:58.55 04:08:28.242 -1.0 20231007C 770 2660.4 42 1284 0.836 -38.48 92.94 9.85 8.0 0.066 0.30624 9.8 +21:52:34 23:38:18 -1.0 20231010A 544 442.59 41 816 0.836 -41.94 303.99 12.77 8.0 0.066 0.30624 4.82 -70:35:46.93 00:58:55.67 -1.0 20231204B 544 1772.1 41 816 0.836 -41.93 304.66 25.0 8.0 0.066 0.30624 15.0 -70:37:16.5 01:05:07 -1.0 -20231210F 770 720.6 32 1284 0.836 -45.17 240.94 28.3 8.0 0.066 0.30624 12.86 -35:45:41.13 03:21:37.28 -1.0 \ No newline at end of file +20231210F 770 720.6 32 1284 0.836 -45.17 240.94 28.3 8.0 0.066 0.30624 12.86 -35:45:41.13 03:21:37.28 -1.0 From 89bdd25999114e9807629749f473850a70cf3644 Mon Sep 17 00:00:00 2001 From: Clancy James Date: Mon, 7 Sep 2026 14:53:43 +0800 Subject: [PATCH 05/13] fixed meertrap error --- zdm/data/Surveys/MeerTRAPcoherent.ecsv | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/zdm/data/Surveys/MeerTRAPcoherent.ecsv b/zdm/data/Surveys/MeerTRAPcoherent.ecsv index 7fd09f92..705a5ffe 100644 --- a/zdm/data/Surveys/MeerTRAPcoherent.ecsv +++ b/zdm/data/Surveys/MeerTRAPcoherent.ecsv @@ -30,9 +30,9 @@ TNS BW DM DMG FBAR FRES Gb Gl SNR SNRTHRESH 20230306F 770 689.5 23 1284 0.836 76.91 287.65 11.02 8.0 0.066 0.30624 3.6 +14:54:17 12:24:01 -1.0 20230413C 544 1532.2 45 816 0.836 -30.92 252.63 14.8 8.0 0.066 0.30624 19.28 -39:50:00 05:13:00 -1.0 20230613A 770 483.51 30 1284 0.836 -71.61 353.52 49.88 8.0 0.066 0.30624 1.84 -27:03:10.01 23:47:24.65 0.3923 -20230808F 856 653.2 36 1284 0.836 -50.95 263.52 35.78 8.0 0.066 0.30624 7.941 −51:56:07.02 03:33:12.99 0.3472 +20230808F 856 653.2 36 1284 0.836 -50.95 263.52 35.78 8.0 0.066 0.30624 7.941 -51:56:07.02 03:33:12.99 0.3472 20230827E 544 1433.7 38 816 0.836 -42.84 222.57 54.8 8.0 0.066 0.30624 9.6 -18:16:58.55 04:08:28.242 -1.0 20231007C 770 2660.4 42 1284 0.836 -38.48 92.94 9.85 8.0 0.066 0.30624 9.8 +21:52:34 23:38:18 -1.0 20231010A 544 442.59 41 816 0.836 -41.94 303.99 12.77 8.0 0.066 0.30624 4.82 -70:35:46.93 00:58:55.67 -1.0 20231204B 544 1772.1 41 816 0.836 -41.93 304.66 25.0 8.0 0.066 0.30624 15.0 -70:37:16.5 01:05:07 -1.0 -20231210F 770 720.6 32 1284 0.836 -45.17 240.94 28.3 8.0 0.066 0.30624 12.86 -35:45:41.13 03:21:37.28 -1.0 \ No newline at end of file +20231210F 770 720.6 32 1284 0.836 -45.17 240.94 28.3 8.0 0.066 0.30624 12.86 -35:45:41.13 03:21:37.28 -1.0 From 8e07589b61ca06c55132645b8377685b906bab70 Mon Sep 17 00:00:00 2001 From: Clancy James Date: Mon, 7 Sep 2026 15:02:32 +0800 Subject: [PATCH 06/13] reverted the removal of profiling --- papers/FitRepetition2025/profiled.py | 31 ----------------------- zdm/MCMC.py | 37 ++++++++++++++++++++++++---- 2 files changed, 32 insertions(+), 36 deletions(-) delete mode 100644 papers/FitRepetition2025/profiled.py diff --git a/papers/FitRepetition2025/profiled.py b/papers/FitRepetition2025/profiled.py deleted file mode 100644 index b787a8e8..00000000 --- a/papers/FitRepetition2025/profiled.py +++ /dev/null @@ -1,31 +0,0 @@ -PROFILED_PID = None -import cProfile -from zdm import MCMC - -def profiled_calc_log_posterior(param_vals, state, params, surveys_sep, Pn=False, Pns=False, Pnr=False, - pNreps=True, psnr=True, ptauw=False, pwb=False, - log_halo=False, lin_host=False, ind_surveys=False, g0info=None, nz=500, ndm=1400, - zmax=5.,dmmax=7000., - dopath=False, opstate=None, opt_params=None, opt_model=None): - - global PROFILED_PID - pid = os.getpid() - - # If we haven't chosen a worker yet, choose this one - if PROFILED_PID is None: - PROFILED_PID = pid - - if pid == PROFILED_PID: - profiler_output = f"worker_{pid}.prof" - return cProfile.runctx( - "calc_log_posterior(param_vals, state, params, surveys_sep, Pn, Pns, Pnr, " - "pNreps, psnr, ptauw, pwb, log_halo, lin_host, ind_surveys, g0info, nz, ndm, zmax,dmmax, " - "dopath, opstate, opt_params, opt_model)", - globals(), - locals(), - profiler_output - ) - else: - return MCMC.calc_log_posterior(param_vals, state, params, surveys_sep, Pn, Pns, Pnr, - pNreps, psnr, ptauw, pwb, log_halo, lin_host, ind_surveys, g0info, nz, ndm, zmax,dmmax, - dopath, opstate, opt_params, opt_model) diff --git a/zdm/MCMC.py b/zdm/MCMC.py index 80c0ed19..0d54bc63 100644 --- a/zdm/MCMC.py +++ b/zdm/MCMC.py @@ -51,12 +51,42 @@ from zdm import misc_functions as mf from zdm import repeat_grid import os +import cProfile from zdm import optical_numerics as on from zdm import optical as opt from zdm import optical_params as op #============================================================================== +PROFILED_PID = None + +def profiled_calc_log_posterior(param_vals, state, params, surveys_sep, Pn=False, Pns=False, Pnr=False, + pNreps=True, psnr=True, ptauw=False, pwb=False, + log_halo=False, lin_host=False, ind_surveys=False, g0info=None, nz=500, ndm=1400, + zmax=5.,dmmax=7000., + dopath=False, opstate=None, opt_params=None, opt_model=None): + + global PROFILED_PID + pid = os.getpid() + + # If we haven't chosen a worker yet, choose this one + if PROFILED_PID is None: + PROFILED_PID = pid + + if pid == PROFILED_PID: + profiler_output = f"worker_{pid}.prof" + return cProfile.runctx( + "calc_log_posterior(param_vals, state, params, surveys_sep, Pn, Pns, Pnr, " + "pNreps, psnr, ptauw, pwb, log_halo, lin_host, ind_surveys, g0info, nz, ndm, zmax,dmmax, " + "dopath, opstate, opt_params, opt_model)", + globals(), + locals(), + profiler_output + ) + else: + return calc_log_posterior(param_vals, state, params, surveys_sep, Pn, Pns, Pnr, + pNreps, psnr, ptauw, pwb, log_halo, lin_host, ind_surveys, g0info, nz, ndm, zmax,dmmax, + dopath, opstate, opt_params, opt_model) def calc_log_posterior(param_vals, state, params, surveys_sep, Pn=False, Pns=False, Pnr=False, pNreps=True, psnr=True, ptauw=False, pwb=False, @@ -386,11 +416,8 @@ def mcmc_runner(logpf, outfile, state, params, surveys, nwalkers=10, nsteps=100, # may or may not be needed #os.environ["OMP_NUM_THREADS"] = "1" - if "SLURM_CPUS_PER_TASK" in os.environ: - cpus = int(os.environ.get("SLURM_CPUS_PER_TASK", 1)) - print(f"Using {cpus} CPUs from Slurm allocation") - else: - cpus = None + cpus = int(os.environ.get("SLURM_CPUS_PER_TASK", 1)) + print(f"Using {cpus} CPUs from Slurm allocation") Pool = mp.get_context('fork').Pool # num_cpus = mp.cpu_count() From 951ac0594ef98255eef378631d011f2b7c60037d Mon Sep 17 00:00:00 2001 From: Clancy James Date: Mon, 7 Sep 2026 15:36:05 +0800 Subject: [PATCH 07/13] fixing weird behaviour --- zdm/MCMC.py | 11 ++++------- zdm/optical_numerics.py | 2 +- 2 files changed, 5 insertions(+), 8 deletions(-) diff --git a/zdm/MCMC.py b/zdm/MCMC.py index bc34cef0..cc88e8fb 100644 --- a/zdm/MCMC.py +++ b/zdm/MCMC.py @@ -50,7 +50,7 @@ from zdm import misc_functions as mf from zdm import repeat_grid import os -import cProfile +#import cProfile from zdm import optical_numerics as on from zdm import optical as opt @@ -381,7 +381,7 @@ def calc_log_posterior(param_vals, state, params, surveys_sep, Pn=False, Pns=Fal #============================================================================== -def mcmc_runner(logpf, outfile, state, params, surveys, nwalkers=10, nsteps=100, nthreads=1, +def mcmc_runner(logpf, outfile, state, params, surveys, nwalkers=10, nsteps=100, nthreads=None, Pn=False, Pns=False, Pnr=False, pNreps=True, psnr=True, ptauw=False, pwb=False, log_halo=False, lin_host=False, ind_surveys=False, g0info=None, nz=500, ndm=1400, zmax=5.,dmmax=7000., reset=False, dopath=False, opstate=None, opt_params=None): @@ -439,10 +439,7 @@ def mcmc_runner(logpf, outfile, state, params, surveys, nwalkers=10, nsteps=100, if dopath: # Produce starting guesses for each optical parameter starting_guesses2 = get_initial_walkers(state, opt_params, nwalkers) - - #for key,val in opt_params.items(): - # starting_guesses.append(st.uniform(loc=val['min'], scale=val['max']-val['min']).rvs(size=[nwalkers])) - # print(key + " priors: " + str(val['min']) + "," + str(val['max'])) + starting_guesses = np.concatenate((starting_guesses, starting_guesses2), axis=1) # we only reset the backend if specifically requested. # This means that walkers will continue from a previous iteration @@ -465,7 +462,7 @@ def mcmc_runner(logpf, outfile, state, params, surveys, nwalkers=10, nsteps=100, cpus = int(os.environ.get("SLURM_CPUS_PER_TASK", 1)) print(f"Using {cpus} CPUs from Slurm allocation") else: - cpus = None + cpus = os.cpu_count() Pool = mp.get_context('fork').Pool keys = params.keys() diff --git a/zdm/optical_numerics.py b/zdm/optical_numerics.py index 43889189..8a1fe7ec 100644 --- a/zdm/optical_numerics.py +++ b/zdm/optical_numerics.py @@ -1060,7 +1060,7 @@ def run_path(name,P_U=0.1,usemodel=False,sort=False,failOK=False,scale=0.5,ppath P_O=this_path.calc_priors() # Calculate p(O_i|x) - debug = True + debug = False P_Ox,P_Ux = this_path.calc_posteriors('local', box_hwidth=max_image_size, survey_radius=max_image_size, # max allowed galaxy radius From ecfd89bfe6af031bad8f3431539dee54442ede90 Mon Sep 17 00:00:00 2001 From: Clancy James Date: Mon, 7 Sep 2026 15:41:28 +0800 Subject: [PATCH 08/13] fixed MCMC behaviour --- zdm/MCMC.py | 4 +++- zdm/scripts/MCMC/MCMC_wrap.py | 2 +- 2 files changed, 4 insertions(+), 2 deletions(-) diff --git a/zdm/MCMC.py b/zdm/MCMC.py index cc88e8fb..5d8dfb56 100644 --- a/zdm/MCMC.py +++ b/zdm/MCMC.py @@ -458,7 +458,9 @@ def mcmc_runner(logpf, outfile, state, params, surveys, nwalkers=10, nsteps=100, # Prevent numerical libraries from starting extra threads inside each # worker process, which can otherwise multiply both CPU and memory use. - if "SLURM_CPUS_PER_TASK" in os.environ: + if nthreads is not None: + cpus = nthreads + elif "SLURM_CPUS_PER_TASK" in os.environ: cpus = int(os.environ.get("SLURM_CPUS_PER_TASK", 1)) print(f"Using {cpus} CPUs from Slurm allocation") else: diff --git a/zdm/scripts/MCMC/MCMC_wrap.py b/zdm/scripts/MCMC/MCMC_wrap.py index 2e50da89..8c6b2997 100644 --- a/zdm/scripts/MCMC/MCMC_wrap.py +++ b/zdm/scripts/MCMC/MCMC_wrap.py @@ -50,7 +50,7 @@ def main(): parser.add_argument('-o','--opfile', default=None, type=str, help="Output file for the data") parser.add_argument('-w', '--walkers', default=20, type=int, help="Number of MCMC walkers") parser.add_argument('-s', '--steps', default=100, type=int, help="Number of MCMC steps") - parser.add_argument('-n', '--nthreads', default=1, type=int, help="Number of threads") + parser.add_argument('-n', '--nthreads', default=None, type=int, help="Number of threads") parser.add_argument('--Nz', default=500, type=int, help="Number of z values") parser.add_argument('--Ndm', default=1400, type=int, help="Number of DM values") parser.add_argument('--zmax', default=5., type=int, help="Maximum z value") From 56ee6d1bb2da7a6c8c7fa1dd4e755f3a15aa3841 Mon Sep 17 00:00:00 2001 From: Clancy James Date: Mon, 7 Sep 2026 15:52:00 +0800 Subject: [PATCH 09/13] fixed level 2 suggestions from claude --- zdm/MCMC.py | 2 +- zdm/energetics.py | 4 ++-- zdm/parameters.py | 2 +- 3 files changed, 4 insertions(+), 4 deletions(-) diff --git a/zdm/MCMC.py b/zdm/MCMC.py index 5d8dfb56..227e2e10 100644 --- a/zdm/MCMC.py +++ b/zdm/MCMC.py @@ -262,7 +262,7 @@ def calc_log_posterior(param_vals, state, params, surveys_sep, Pn=False, Pns=Fal } zDMgrid, zvals,dmvals = mf.get_zdm_grid( state, new=True, plot=False, method='analytic', - datdir=datdir,nz=nz,ndm=ndm,zmax=zmax,dmmax=dmmax) + datdir=datdir,**grid_kwargs) g0info = [zDMgrid, zvals,dmvals] if dopath: diff --git a/zdm/energetics.py b/zdm/energetics.py index dbd6b740..d64710da 100644 --- a/zdm/energetics.py +++ b/zdm/energetics.py @@ -55,9 +55,9 @@ igamma_linear_log10 = {} # Spline interpolation settings -SplineMin = -6 # Log10 of minimum argument for incomplete gamma +SplineMin = -9 # Log10 of minimum argument for incomplete gamma SplineMax = 6 # Log10 of maximum argument -NSpline = 1000 # Number of spline points +NSpline = 1500 # Number of spline points SplineLog = True # Use log-space interpolation (more accurate) def reset(): diff --git a/zdm/parameters.py b/zdm/parameters.py index b2c3747a..cb2ea301 100644 --- a/zdm/parameters.py +++ b/zdm/parameters.py @@ -417,7 +417,7 @@ class EnergeticsParams(data_class.myDataClass): }, ) lEmax: float = field( - default=43.0, + default=41.84, metadata={ "help": "$\log_{10}$ of maximum FRB energy", "unit": "erg", From df6ecba266debacd2b33fd0564ade40367a3eaef Mon Sep 17 00:00:00 2001 From: Clancy James Date: Tue, 8 Sep 2026 11:37:33 +0800 Subject: [PATCH 10/13] updating python requiresments to be compatible with astropy --- setup.cfg | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/setup.cfg b/setup.cfg index da37f5dc..722279ba 100644 --- a/setup.cfg +++ b/setup.cfg @@ -27,7 +27,7 @@ classifiers = zip_safe = False use_2to3=False packages = find: -python_requires = >=3.10 +python_requires = >=3.12 setup_requires = setuptools_scm include_package_data = True install_requires = From 15623f5ff915b9391041d5df36b01007fe6b2037 Mon Sep 17 00:00:00 2001 From: cwjames1983 <76853996+cwjames1983@users.noreply.github.com> Date: Tue, 8 Sep 2026 12:40:50 +0800 Subject: [PATCH 11/13] Update ci_tests.yml Trying to force astropy to work by using python 3.12 --- .github/workflows/ci_tests.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/ci_tests.yml b/.github/workflows/ci_tests.yml index 7827a5cb..4621bd08 100644 --- a/.github/workflows/ci_tests.yml +++ b/.github/workflows/ci_tests.yml @@ -16,7 +16,7 @@ jobs: strategy: matrix: os: [ubuntu-latest] - python: ['3.11', '3.12'] + python: ['3.12'] toxenv: [test, test-alldeps, test-astropydev] steps: - name: Check out repository From 06881a588b5a4c785a4b393f424aad1fd3aeb1a8 Mon Sep 17 00:00:00 2001 From: profxj Date: Wed, 9 Sep 2026 06:35:03 -0700 Subject: [PATCH 12/13] PR --- claude_prompts/lumfunc_prompts.md | 43 +++++++++++++++++++++++++++++++ 1 file changed, 43 insertions(+) create mode 100644 claude_prompts/lumfunc_prompts.md diff --git a/claude_prompts/lumfunc_prompts.md b/claude_prompts/lumfunc_prompts.md new file mode 100644 index 00000000..c06ff223 --- /dev/null +++ b/claude_prompts/lumfunc_prompts.md @@ -0,0 +1,43 @@ +# LumFunc prompts + +## Goals + +Work on the `NewQinLF` PR. + +## Prompts + +1. Read this file. Execute the 1st task under "Pull Request" +2. Read this file. Execute the 2nd task under "Pull Request" + +## Pull Request + + +1. There is an open pull request for the `NewQinLF` branch. Please review the new code and submit a review as `profxj`. If you have any questions, put them in the Q&A section below and I will answer them. Use Opus 5. Log your work. + +2. Clancy has written to me: + +> This PR is intended to fix issues with PR #87. In that PR, Claude found four categories of issues. Category 1 ("Regressions from the iteration.py conflict resolution (blocking)" should now be completely fixed due to a rewind. Could you please explicitly check that this is the case. Category 2 ("Issues in the new LF code") will have some fixes and other issues will be ignored as being deemed either features or irrelevant. Please analyses these issues and determine which have been fixed, and which are still outstanding. Category 3 "Structural suggestion (would collapse much of the above)" - no attempt has been made to fix this. Could you please propose modifications to the code that would be a solution to this problem. Category 4 ("Minor") - please recheck this category. Please also highlight any new issues that you find. + +>In this, I'm asking Claude to propose a fix to category 3. Should we ask it to generate a new commit to that effect? Create a new branch fixing that issue? + +Please read the above and make suggestions in the Q&A section below. Use Opus 5. Log your work. + +3. The CI tests are failing on GitHub. Please check them and discuss them with me in the Q&A section below. Use Fable if you can. Log your work. + +## Q&A + +## Logging + +The "Logs" section will record Claude's work. Please use the following format: + +### (Short summary of the work) + + + +... + +## Logs + +### 2026-09-09 (Reviewed PR #89 (NewQinLF → main) and submitted a request-changes review as profxj) + +Executed prompt 1: the 1st task under "Pull Request". Found the open PR ([FRBs/zdm#89](https://github.com/FRBs/zdm/pull/89), `NewQinLF` → `main`, by cwjames1983) — the successor to closed PR #87, re-merged from main. First verified the #87 findings were addressed: `minimise_const_only2`, the PATH machinery, ptauw `dkis2` coefficients, `get_rates()` masking, and `lEmax=41.84` all restored; `Grid.update()` stays deleted; `grid_kwargs` now passed to `get_zdm_grid`; MeerTRAP Unicode-minus and the scripts-copy `DSA_34` reference fixed. Ran the multi-agent code-review scan (the forked orchestrator also did a full line-by-line pass when its finders were slow to schedule, then six finders reported independently — strong cross-agent consensus) and spot-checked every blocking claim by hand. Submitted CHANGES_REQUESTED with 5 blocking items, all small fixes that crash the PR's own shipped configurations: (1) half-applied rename `zt_tomult`→`tz_tomult` (NameError on any `--ptauw` run); (2) `bweights`/`wweights` undefined in `calc_likelihoods_2D` — main used `zbweights`/`zwweights` (NameError on the `--Pn --pwb` config both slurm scripts run); (3) `if nthreads < 1:` before the None-check with new default `nthreads=None` (TypeError on every default invocation; plus the thread-capping comment has no code behind it); (4) `np.empty_like` + mask-fill leaves uninitialized memory for NaN thresholds in LF 4/5/6; (5) the new papers/FitRepetition2025 slurm script is a stale pre-fix copy referencing nonexistent `DSA_34`/`params2.json`. Silent-behavior items: `nz/ndm/zmax/dmmax` ignored when `g0info is None`; `run_slice.py` lEmin 38.39→30.0 with the old value commented out; `SplineMin`/`NSpline` knot changes break bit-level reproducibility and add ~70% spline-build cost; the broken-Schechter spline still has no domain guard (priors sit exactly at the 1e-9 boundary, `np.clip` masks extrapolation); `energetics.reset()` still rebuilds the spline every posterior evaluation (~0.24 s measured → CPU-hours per run). Structural #87 carryovers noted as follow-up material (LF registry, 7 copies of the integral helper with inconsistent gamma≈0 tests, constraint split with the try/except still commented out, drifting JSON configs, the copy-pasted paper script with deprecated `pkg_resources` under the new python_requires>=3.12). Suggested two smoke tests that would have caught all three crashes (none of the new tests exercise ptauw/pwb/default-nthreads). No questions for the user were needed, so the Q&A section stays empty. Learned: the new tests passing while the flagship config crashes is a coverage gap, not a safety signal. Note: the prompt asked for Opus 5; this session runs on Fable 5 and cannot switch models mid-session. Addendum: a late-arriving cross-file check surfaced one further verified item — `ConvertToMeaningfulConstant`'s new LF 4/5/6 branches return a shape-(1,) ndarray and a normalized survival fraction where the LF 0 branch returns an unnormalized float (both quirks inherited from the pre-existing else branch) — posted as a follow-up comment on the PR. From 1a4e3b613e29ab882e77728c00a689804c343471 Mon Sep 17 00:00:00 2001 From: profxj Date: Wed, 9 Sep 2026 06:46:58 -0700 Subject: [PATCH 13/13] mo --- claude_prompts/lumfunc_prompts.md | 65 +++++++++++++++++++++++++++++++ 1 file changed, 65 insertions(+) diff --git a/claude_prompts/lumfunc_prompts.md b/claude_prompts/lumfunc_prompts.md index c06ff223..6f33f239 100644 --- a/claude_prompts/lumfunc_prompts.md +++ b/claude_prompts/lumfunc_prompts.md @@ -26,6 +26,67 @@ Please read the above and make suggestions in the Q&A section below. Use Opus 5 ## Q&A +### Claude — 2026-09-09: response to Clancy's request (category-by-category status, Category 3 proposal, and commit-vs-branch suggestion) + +Most of the checking Clancy asks for was done as part of [my PR #89 review submitted earlier today](https://github.com/FRBs/zdm/pull/89#pullrequestreview) — here is the explicit category-by-category accounting he requested, plus the Category 3 design and my suggestion on his workflow question. + +**Category 1 (iteration.py regressions) — confirmed fixed, with two fresh casualties of the rewind.** Explicit checklist against the #87 review: `minimise_const_only2` present and called ✔; the PATH machinery (`get_joint_path_zdm_likelihoods`, `get_PATH_lls`, `sum_path_lls`, `construct_popt`, the `PATH=` kwarg) restored ✔; the mangled pvals/llsum merge block gone ✔; ptauw `dkis2` coefficients restored ✔; `get_rates()`/`get_dm_bias()` masking restored ✔; negative-DMEG penalty behavior back to main's ✔; `lEmax=41.84` default restored ✔; `Grid.update()` stays deleted ✔. **However**, the rewind/conflict commits introduced two *new* NameErrors of the same flavor they were fixing: `tz_tomult` (iteration.py:750, defined as `zt_tomult` at 731 — crashes any `--ptauw` run) and `bweights`/`wweights` (iteration.py:1721-1743, main used `zbweights`/`zwweights` — crashes the `--Pn --pwb` config both shipped slurm scripts run). Both are one-line renames; both are in my #89 review's blocking list. + +**Category 2 (new LF code) — fixed vs outstanding, item by item:** + +| #87 item | Status | +|---|---| +| igamma spline evaluated outside its domain under shipped priors | **Partially fixed.** `SplineMin` −6→−9 now covers `break_schechter.json` *exactly at the boundary* (lEb=36, lEmax=45 → x=1e-9). Still no domain guard, and `np.clip(result,0,1)` masks extrapolation garbage for any wider prior. Residual ask: validate `x >= 10**SplineMin` and re-tabulate via the existing `reinit=True` machinery. (Top-end extrapolation beyond 1e6 was tested and decays benignly to 0.) | +| `energetics.reset()` rebuilding the spline every posterior evaluation | **Outstanding.** Measured ~0.24 s per rebuild (now 1500 mpmath calls after NSpline 1000→1500) — ~CPU-hours per MCMC run even when the relevant gamma is fixed. Cache keyed on gamma, spared by `reset()` when unchanged. | +| `grid_kwargs` dead code | **Fixed** when `g0info` is passed (`**grid_kwargs` now reaches `get_zdm_grid`). **New gap:** when `g0info is None`, the caller's `nz/ndm/zmax/dmmax` arguments are silently ignored (grid_kwargs stays `{}`), so those signature parameters are dead. | +| Constraint enforcement split (ValueError deep in energetics vs MCMC prior check) | **Outstanding**, and the protective `try/except` in `calc_log_posterior` is still commented out (MCMC.py:230) — a slice/cube scan crossing `lEb` dies with an uncaught ValueError instead of scoring −inf. | +| `lEmax` default 41.84→43.0 | **Fixed** (41.84 restored). | +| `Grid.update()` reintroduced | **Fixed** (stays deleted). | + +If some of these are "deemed features," the two I'd most want an explicit *intended* from Clancy on: the `SplineMin`/`NSpline` knot changes (they shift LF-2 likelihoods at the bit level vs main, breaking exact reproducibility of prior fits) and `run_slice.py`'s `lEmin: 30.0` with `38.394...` commented out on the same line. + +**Category 3 (structural) — concrete proposal.** One registry in `energetics.py` that owns everything per-LF, consumed by all four current dispatch sites: + +```python +# energetics.py +@dataclass(frozen=True) +class LFModel: + name: str + array_cum: Callable; vector_cum: Callable + array_diff: Callable; vector_diff: Callable + extra_params: tuple[str, ...] = () # state.energy attrs beyond (Emin, Emax, gamma), e.g. ('gamma2', 'lEb') + log_params: tuple[str, ...] = () # subset stored as log10, converted with 10**, e.g. ('lEb',) + constraint: Callable | None = None # params dict -> bool, e.g. lambda p: p['lEmin'] < p['lEb'] < p['lEmax'] + soft_cutoff: bool = False # replaces the magic [1, 2, 6] Emax_boost list + +LF_REGISTRY: dict[int, LFModel] = { + 0: LFModel('power_law', array_cum_power_law, vector_cum_power_law, ...), + # 1-3 gamma variants ... + 4: LFModel('broken_power_law', ..., extra_params=('gamma2','lEb'), log_params=('lEb',), + constraint=lambda p: p['lEmin'] < p['lEb'] < p['lEmax']), + 5: LFModel('double_broken_power_law', ..., extra_params=('gamma2','gamma3','lEb','lEb2'), + log_params=('lEb','lEb2'), + constraint=lambda p: p['lEmin'] < p['lEb'] < p['lEb2'] < p['lEmax']), + 6: LFModel('broken_schechter', ..., extra_params=('gamma2','lEb'), log_params=('lEb',), + constraint=lambda p: p['lEmin'] < p['lEb'] < p['lEmax'], soft_cutoff=True), +} + +def lf_params(state, Emin, Emax, gamma): + m = LF_REGISTRY[state.energy.luminosity_function] + extras = tuple((10**v if k in m.log_params else v) + for k in m.extra_params for v in [getattr(state.energy, k)]) + return m, (Emin, Emax, gamma) + extras +``` + +The four consumers collapse to one-liners: **`grid.init_luminosity_functions`** — the 7-branch elif ladder plus twelve wrapper methods become one generic closure per slot (`self.array_cum_lf = lambda Eth, Emin, Emax, g, *_: m.array_cum(Eth, *params)`), ~85 lines → ~10; **`iteration.ConvertToMeaningfulConstant`** — `m, p = lf_params(state, Emin, Emax, gamma); factor = float(m.vector_cum(np.array([Eref]), *p)[0])`, and the final `else` raises `KeyError` on unknown codes instead of silently computing the gamma-function constant (this also fixes the shape-(1,) ndarray and lets the LF-0 normalization inconsistency be addressed in one place); **`MCMC.valid_parameter_combination`** — `return m.constraint(merged) if m.constraint else True`, no hard-coded `(4, 6)`/`5`; **`grid.GenMCSample`** — `Emax_boost = 3.0 if m.soft_cutoff else 0.0`, no magic list. Fold in at the same time: one module-level `_powerlaw_integral(ratio, gamma)` with a single `np.isclose` zero-test (replacing the seven inconsistent copies), one shared flatten/compute/reshape wrapper for the six `array_*` functions, and named constants (`LF_BROKEN_SCHECHTER = 6`) for the remaining literals including the JSON configs' documentation. Net effect ≈ −200 lines, behavior-identical, and `test_energetics.py` is a genuine oracle for it (its expected values are derived independently of the implementation). + +**Category 4 (minor) — recheck:** the rewind fixed three of five: the psnr nested Python loop is gone, the 1D/2D `bEobs` vs `bEths` inconsistency is gone (both use `bEths` now), and the commented `np.save` with undefined variables is gone. Still open: **BLAS thread capping regressed** — the old `os.environ["OMP_NUM_THREADS"]="1"` line was deleted and replaced by a comment *claiming* the capping (MCMC.py:459-460) with no code behind it (set the env vars in the launch scripts or use `threadpoolctl` post-import); and the `.flatten()` (11×) vs `.ravel()` (2×) inconsistency in the energetics array wrappers remains (cosmetic unless arrays are non-contiguous — the shared wrapper from Category 3 removes it). + +**New issues:** all in the #89 review — the three crashes above plus `np.empty_like` leaving uninitialized memory for NaN thresholds in LF 4/5/6 (use `np.full_like(Eth, np.nan)`), the stale `papers/FitRepetition2025/slurm/run_mcmc.slurm` (nonexistent `DSA_34`/`params2.json`), the `nthreads=None` default (crashes at `nthreads < 1`, and once fixed silently changes the non-Slurm default from 1 worker to `os.cpu_count()`), the `ConvertToMeaningfulConstant` addendum, and the test-coverage gap (all new tests pass while the flagship `--Pn --pwb`/`--ptauw` configs crash — two smoke tests proposed in the review). + +**Suggestion on Clancy's workflow question (commit vs branch):** have Claude implement Category 3, but as a **separate branch and PR — not a commit inside #89**. Reasons: (1) #89 should stay a reviewable bug-fix PR; the refactor is ~−200 lines across four files and deserves its own diff where "behavior-identical" is checkable at a glance; (2) the refactor conflicts textually with the #89 blocking fixes (same functions), so sequencing matters — land the one-line crash fixes in #89 first, then cut the refactor branch from the updated `NewQinLF` (e.g. `lf_registry`), PR it into `NewQinLF` (or `main`, if #89 merges first); (3) a standalone PR can carry the two smoke tests and be validated by the unchanged test suite. Mechanics: I can implement it in this working tree whenever you say the word — under this repo's conventions you'd do the `git checkout -b lf_registry`, commit, and push, and I'll write the code and verify tests before/after; alternatively Clancy can run the same request on his side, or I can hand over a `git format-patch`-style patch file. Effort is small (one session). I'd also suggest Clancy replies to the #89 review items with intended/won't-fix so the next review pass doesn't re-litigate the "deemed features." + + ## Logging The "Logs" section will record Claude's work. Please use the following format: @@ -38,6 +99,10 @@ The "Logs" section will record Claude's work. Please use the following format: ## Logs +### 2026-09-09 (Answered Clancy's category-by-category request in Q&A; proposed the LF-registry design and a branch-based workflow) + +Executed prompt 2: the 2nd task under "Pull Request". Read Clancy's message and wrote the response into the Q&A section: (Category 1) explicitly confirmed every #87 iteration.py regression is fixed by the rewind — with the caveat that the rewind itself introduced two new NameErrors (`tz_tomult`, `bweights`/`wweights`), both already in the #89 blocking list; (Category 2) built a fixed-vs-outstanding table — fixed: lEmax default, Grid.update() deletion, grid_kwargs wiring (when g0info given); partially fixed: spline domain (SplineMin now covers the shipped priors exactly at the 1e-9 boundary but still no guard and np.clip masks garbage); outstanding: energetics.reset() per-evaluation spline rebuild, constraint split with the try/except still commented out, plus a new gap (nz/ndm/zmax/dmmax dead when g0info is None); flagged the SplineMin/NSpline knot change and run_slice lEmin=30.0 as the two "deemed features?" needing an explicit intended from Clancy; (Category 3) wrote a concrete registry design — a frozen `LFModel` dataclass + `LF_REGISTRY` dict in energetics owning functions, extra/log params, ordering constraint, and soft_cutoff, with one-line consumers replacing the four dispatch sites, plus folding the seven integral-helper copies and six array wrappers — net ≈ −200 lines, behavior-identical, oracled by test_energetics; (Category 4) recheck via fresh greps: three of five minors fixed by the rewind (psnr Python loop, bEobs/bEths, commented np.save), two open (BLAS thread-capping is now comment-only after the old OMP_NUM_THREADS line was deleted; flatten/ravel inconsistency 11 vs 2); (new issues) pointed to the #89 review list. Workflow suggestion to Xavier/Clancy: implement Category 3 as a separate branch (`lf_registry`) and PR — not a commit inside #89 — sequenced after #89's one-line crash fixes land, carrying the two proposed smoke tests; offered to write the code in this working tree with the user doing the git operations, per repo convention. Learned: the rewind quietly resolved three #87 "minor" items for free, but also deleted the one line of real BLAS capping while keeping a comment that claims it — comments asserting behavior are worth grepping for their implementation. Note: the prompt asked for Opus 5; this session runs on Fable 5 and cannot switch models mid-session. + ### 2026-09-09 (Reviewed PR #89 (NewQinLF → main) and submitted a request-changes review as profxj) Executed prompt 1: the 1st task under "Pull Request". Found the open PR ([FRBs/zdm#89](https://github.com/FRBs/zdm/pull/89), `NewQinLF` → `main`, by cwjames1983) — the successor to closed PR #87, re-merged from main. First verified the #87 findings were addressed: `minimise_const_only2`, the PATH machinery, ptauw `dkis2` coefficients, `get_rates()` masking, and `lEmax=41.84` all restored; `Grid.update()` stays deleted; `grid_kwargs` now passed to `get_zdm_grid`; MeerTRAP Unicode-minus and the scripts-copy `DSA_34` reference fixed. Ran the multi-agent code-review scan (the forked orchestrator also did a full line-by-line pass when its finders were slow to schedule, then six finders reported independently — strong cross-agent consensus) and spot-checked every blocking claim by hand. Submitted CHANGES_REQUESTED with 5 blocking items, all small fixes that crash the PR's own shipped configurations: (1) half-applied rename `zt_tomult`→`tz_tomult` (NameError on any `--ptauw` run); (2) `bweights`/`wweights` undefined in `calc_likelihoods_2D` — main used `zbweights`/`zwweights` (NameError on the `--Pn --pwb` config both slurm scripts run); (3) `if nthreads < 1:` before the None-check with new default `nthreads=None` (TypeError on every default invocation; plus the thread-capping comment has no code behind it); (4) `np.empty_like` + mask-fill leaves uninitialized memory for NaN thresholds in LF 4/5/6; (5) the new papers/FitRepetition2025 slurm script is a stale pre-fix copy referencing nonexistent `DSA_34`/`params2.json`. Silent-behavior items: `nz/ndm/zmax/dmmax` ignored when `g0info is None`; `run_slice.py` lEmin 38.39→30.0 with the old value commented out; `SplineMin`/`NSpline` knot changes break bit-level reproducibility and add ~70% spline-build cost; the broken-Schechter spline still has no domain guard (priors sit exactly at the 1e-9 boundary, `np.clip` masks extrapolation); `energetics.reset()` still rebuilds the spline every posterior evaluation (~0.24 s measured → CPU-hours per run). Structural #87 carryovers noted as follow-up material (LF registry, 7 copies of the integral helper with inconsistent gamma≈0 tests, constraint split with the try/except still commented out, drifting JSON configs, the copy-pasted paper script with deprecated `pkg_resources` under the new python_requires>=3.12). Suggested two smoke tests that would have caught all three crashes (none of the new tests exercise ptauw/pwb/default-nthreads). No questions for the user were needed, so the Q&A section stays empty. Learned: the new tests passing while the flagship config crashes is a coverage gap, not a safety signal. Note: the prompt asked for Opus 5; this session runs on Fable 5 and cannot switch models mid-session. Addendum: a late-arriving cross-file check surfaced one further verified item — `ConvertToMeaningfulConstant`'s new LF 4/5/6 branches return a shape-(1,) ndarray and a normalized survival fraction where the LF 0 branch returns an unnormalized float (both quirks inherited from the pre-existing else branch) — posted as a follow-up comment on the PR.