From 2f0a967517c3ab060db679771f1f8a3da9da48ce Mon Sep 17 00:00:00 2001 From: Lenard Dome <23503601+lenarddome@users.noreply.github.com> Date: Thu, 24 Sep 2026 17:35:22 +0200 Subject: [PATCH 01/34] perf(generators): share priors between copies, evaluate them in closed form Also fixes cpm.hierarchical discarding its population priors after the first evaluation. Co-Authored-By: Claude Opus 5.5 (1M context) --- CHANGELOG.md | 5 + benchmarks/cases.py | 111 ++++++++++++++++++ benchmarks/results/main.csv | 6 + benchmarks/results/phase0.csv | 6 + benchmarks/run.py | 115 +++++++++++++++++++ cpm/generators/_fast_priors.py | 147 ++++++++++++++++++++++++ cpm/generators/parameters.py | 24 +++- cpm/generators/wrapper.py | 52 ++++++++- cpm/hierarchical/empirical.py | 4 +- cpm/hierarchical/variational.py | 4 +- test/generators/test_fast_priors.py | 172 ++++++++++++++++++++++++++++ 11 files changed, 637 insertions(+), 9 deletions(-) create mode 100644 benchmarks/cases.py create mode 100644 benchmarks/results/main.csv create mode 100644 benchmarks/results/phase0.csv create mode 100644 benchmarks/run.py create mode 100644 cpm/generators/_fast_priors.py create mode 100644 test/generators/test_fast_priors.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 45679af..9d58e4f 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -28,9 +28,14 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 - Converted all docstrings from Markdown to reStructuredText (NumPy style), fixed incorrect examples in the docstrings of `RLRW`, `EmpiricalBayes`, `ProspectUtility`, `bin_ratings`, `count_trials` and `convert_to_RLRW`, added an example to `VariationalBayes`, and corrected the column list of `load_bandit_data`. Docstrings with LaTeX are now raw strings, which fixes a `SyntaxWarning` and a corrupted equation in the docstring of `RLRW` - Restricted package discovery to `cpm`, so that files under `docs/` are no longer packaged, and added `docs` and `notebooks` optional dependencies - Significantly improved the performance of `cpm.hierarchical.EmpiricalBayes` and `cpm.hierarchical.VariationalBayes` for long or multi-chain EM runs, by buffering results instead of concatenating on every iteration and vectorising the Hessian inversions +- Made each evaluation of the objective function cheaper for every model, by 0.5-1 ms for 2-6 free parameters (about 0.1 ms per free parameter): `cpm.generators.Value` copies now share their prior distribution instead of deep-copying it, `cpm.generators.Wrapper.reset()` restores parameter values and initial states without copying the priors, and `Value.PDF(log=True)` evaluates the uniform, truncated normal, beta, gamma, truncated exponential and normal priors (including frozen scipy distributions of these families passed in by the user) in closed form instead of through scipy. The priors remain frozen scipy distributions, and the log densities equal scipy's to 1e-12. Models that copy their parameters on every trial gain most: `PTSM2025` runs 17 times faster +- `cpm.generators.Value.update_prior()` now replaces the prior with an updated copy instead of changing it in place, since copies of a `Value` share their prior. Code that changes `value.prior.kwds` directly now changes the prior of every copy of that `Value` +- Added a benchmark suite under `benchmarks/`, which times each part of the objective function for every built-in application (`python benchmarks/run.py`) ### Fixed +- Fixed `cpm.hierarchical.EmpiricalBayes` and `cpm.hierarchical.VariationalBayes` fitting every participant with the model's original priors instead of the estimated population-level priors, apart from the first evaluation of the objective function. `Wrapper.reset()` restored the parameters, including their priors, from a copy made when the model was created, so every prior update was discarded after the first run of the model. Hierarchical fits now use the updated priors throughout, which changes their results +- Fixed the random starting priors of the second and later chains of `cpm.hierarchical.EmpiricalBayes` and `cpm.hierarchical.VariationalBayes` being arrays of length one instead of numbers, which made the log prior an array. It went unnoticed because of the previous issue - Fixed `cpm.optimisation.Bads` emitting a `DeprecationWarning` on every GP fit, which would become an error on Python 3.16, by passing `gp_fixed_mean` to pybads as a NumPy boolean ([#88](https://github.com/DevComPsy/cpm/issues/88)) - Fixed `cpm.optimisation.Bads` failing with `ValueError: setting an array element with a sequence` under NumPy 2 by requiring `gpyreg>=1.2.1` - Fixed `cpm.optimisation.FminBound` raising `TypeError` on SciPy 1.18.0 and later, which removed the `disp` and `iprint` options of L-BFGS-B; they are now passed only where SciPy still accepts them diff --git a/benchmarks/cases.py b/benchmarks/cases.py new file mode 100644 index 0000000..0eaf57a --- /dev/null +++ b/benchmarks/cases.py @@ -0,0 +1,111 @@ +""" +The benchmark cases: every built-in application of cpm, on one participant. + +Each case builds a model for a given backend and returns it together with the +observed data and the parameter values at which the objective is evaluated. +`backend=None` is the per-trial application (the original classes); the other +backends are the session versions. Cases whose session version does not exist +in the installed cpm are skipped, so the same file also benchmarks older +versions of cpm. +""" + +import warnings + +import numpy as np +import pandas as pd + +import cpm +from cpm.datasets import load_bandit_data, load_risky_choices + + +def bandit(): + data = load_bandit_data() + data = data[data.ppt == 1].reset_index(drop=True) + data["observed"] = data["response"] + return data + + +def two_step(n=200, seed=0): + rng = np.random.default_rng(seed) + action = rng.integers(0, 2, n) + common = rng.random(n) < 1.0 # deterministic transitions in the novel task + return pd.DataFrame( + { + "s1": rng.integers(0, 2, n), + "stimuli_first": rng.integers(0, 2, n), + "action": action, + "s2": np.where(common, 1 - action, action), + "reward": rng.integers(0, 10, n) / 9, + "reward_0": rng.integers(0, 10, n) / 9, + "reward_1": rng.integers(0, 10, n) / 9, + "observed": action, + } + ) + + +def risky(): + data = load_risky_choices() + data = data[data.ppt == 1].reset_index(drop=True) + data["observed"] = data["choice"].astype(int) + return data + + +PT_SETTINGS = { + "alpha": [0.8, 1e-2, 5.0], + "lambda_loss": [1.6, 1e-2, 5.0], + "gamma": [0.7, 1e-2, 5.0], + "temperature": [3.0, 1e-2, 15.0], + "beta": [0.9, 0.0, 5.0], + "delta": [0.6, 1e-2, 5.0], + "eta": [0.1, -0.49, 0.49], + "phi_gain": [0.2, -10.0, 10.0], + "phi_loss": [-0.3, -10.0, 10.0], +} + + +def _application(module, name, backend): + """The per-trial class for backend None, else its session version (or None).""" + module = getattr(cpm.applications, module) + if backend is None: + return getattr(module, name) + session = getattr(module, name + "Session", None) + if session is None: + return None + return lambda **kwargs: session(backend=backend, **kwargs) + + +def build(case, backend=None): + """ + The model, observed data and parameter values of a benchmark case. + + Returns None if the installed cpm has no model for this backend. + """ + warnings.simplefilter("ignore") + if case == "RLRW": + cls = _application("reinforcement_learning", "RLRW", backend) + data = bandit() + kwargs = dict(dimensions=4, parameters_settings=[[0.3, 0, 1], [4, 0, 10]]) + elif case == "HybridMBMF": + cls = _application("reinforcement_learning", "HybridMBMF", backend) + data = two_step() + kwargs = dict( + parameters_settings=[ + [2, 0, 5], [0.4, 0, 1], [0.6, 0, 1], [0.5, 0, 1], [0.3, -5, 5], [-0.2, -5, 5] + ] + ) + elif case in ("PTSM", "PTSM1992", "PTSM2025"): + cls = _application("decision_making", case, backend) + data = risky() + kwargs = dict(parameters_settings=PT_SETTINGS) + else: + raise ValueError(case) + if cls is None: + return None + model = cls(data=data, **kwargs) + observed = data["observed"].to_numpy() + x = np.array([model.parameters[k].value for k in model.parameters.free()], dtype=float) + return model, observed, x + + +CASES = ["RLRW", "HybridMBMF", "PTSM", "PTSM1992", "PTSM2025"] +BACKENDS = [None, "python", "numba"] diff --git a/benchmarks/results/main.csv b/benchmarks/results/main.csv new file mode 100644 index 0000000..2ec3a77 --- /dev/null +++ b/benchmarks/results/main.csv @@ -0,0 +1,6 @@ +case,backend,trials,parameters,objective_ms,reset_ms,prior_ms,run_ms,loss_ms,log_posterior,speed_up +RLRW,trial,71,2,8.6474500130862,0.42120006401091814,0.08180004078894854,7.56499997805804,0.04610000178217888,-26.034401210112605,1.0 +HybridMBMF,trial,200,6,25.437099975533783,0.8433500188402832,0.24389999452978373,23.479699972085655,0.05304999649524689,-151.7819350468299,1.0 +PTSM,trial,40,4,6.54054997721687,0.6358500104397535,0.15730003360658884,5.147499963641167,0.04579999949783087,-36.42346848683204,1.0 +PTSM1992,trial,40,6,7.78515002457425,0.8213000255636871,0.22419996093958616,6.095800083130598,0.04350004019215703,-39.43873213610623,1.0 +PTSM2025,trial,40,5,75.79680002527311,0.6942999898456037,0.19024999346584082,74.84180002938956,0.04529999569058418,-33.01065168146936,1.0 diff --git a/benchmarks/results/phase0.csv b/benchmarks/results/phase0.csv new file mode 100644 index 0000000..7efdf3c --- /dev/null +++ b/benchmarks/results/phase0.csv @@ -0,0 +1,6 @@ +case,backend,trials,parameters,objective_ms,reset_ms,prior_ms,run_ms,loss_ms,log_posterior,speed_up +RLRW,trial,71,2,7.531950017437339,0.004799978341907263,0.003999972250312567,7.398999994620681,0.04650000482797623,-26.034401210112605,1.0 +HybridMBMF,trial,200,6,23.840799927711487,0.01279998105019331,0.007400056347250938,23.64180004224181,0.0543000060133636,-151.7819350468299,1.0 +PTSM,trial,40,4,5.247199966106564,0.0073499977588653564,0.005299923941493034,5.067099933512509,0.04069996066391468,-36.42346848683204,1.0 +PTSM1992,trial,40,6,6.09894999070093,0.010300078429281712,0.007400056347250938,5.823099985718727,0.046749948523938656,-39.438732136106225,1.0 +PTSM2025,trial,40,5,4.435600014403462,0.0071999384090304375,0.005500041879713535,4.139299970120192,0.04824995994567871,-33.01065168146936,1.0 diff --git a/benchmarks/run.py b/benchmarks/run.py new file mode 100644 index 0000000..6465dbf --- /dev/null +++ b/benchmarks/run.py @@ -0,0 +1,115 @@ +""" +Where does the time of one evaluation of the objective go, for every built-in +application of cpm and every backend? + + python benchmarks/run.py [--repeats 200] [--label NAME] + +Prints a table and writes benchmarks/results/