Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
35 commits
Select commit Hold shift + click to select a range
2f0a967
perf(generators): share priors between copies, evaluate them in close…
lenarddome Sep 24, 2026
09db79c
feat(generators): add SessionWrapper, for models that run a whole ses…
lenarddome Sep 24, 2026
4b051fa
feat: add numba-compiled session versions of the built-in applications
lenarddome Sep 24, 2026
68c2d9f
perf: speed up cpm.models, per-trial runs, the losses and meta-d'
lenarddome Sep 24, 2026
aadbc63
refactor(applications): one implementation per model, same API
lenarddome Sep 25, 2026
f7eb422
refactor(models): one formula per building block, shared with the com…
lenarddome Sep 28, 2026
eaebd8b
docs: add the numba install option
lenarddome Sep 28, 2026
e145a7c
docs(changelog): bring the numba entries up to date
lenarddome Sep 28, 2026
755cd3f
docs(changelog): shorten the entries of the development version
lenarddome Sep 28, 2026
cc7160d
chore: bump the version to 0.26.0.dev0
lenarddome Sep 28, 2026
2d48bad
bench: time the hierarchical tutorials; correct the compile times in …
lenarddome Sep 28, 2026
9debd33
docs: re-run the hierarchical tutorials
lenarddome Sep 28, 2026
6b31a12
docs: add a News section, with a post on faster fitting in cpm 0.26
lenarddome Sep 28, 2026
8c61db8
fix(hierarchical): seed the starting priors of later chains, within b…
lenarddome Sep 28, 2026
9d275aa
docs(tutorials): point the hierarchical tutorials to RLRW, and re-run…
lenarddome Sep 28, 2026
63ba29d
docs: add a how-to guide on speeding up your own model
lenarddome Sep 28, 2026
fd56207
docs(numba): correct the caching and debugging advice
lenarddome Sep 28, 2026
832a2ff
fix(applications): check indices and activations in the compiled loops
lenarddome Sep 29, 2026
03ffcf0
fix(generators): pickling and invalid priors, and SessionWrapper and …
lenarddome Sep 29, 2026
ee689cd
fix(jit): fall back to plain Python or no caching when numba cannot b…
lenarddome Sep 29, 2026
dcb3b93
chore: fix the wheel tag, Python requirement and classifiers; drop un…
lenarddome Sep 29, 2026
18b44bb
ci: run the tests on Python 3.11-3.14, with and without numba
lenarddome Sep 29, 2026
a1a3e41
docs(changelog): note that Wrapper.reset keeps the current priors
lenarddome Sep 29, 2026
897f97f
chore(benchmarks): remove results that nothing reads
lenarddome Sep 29, 2026
f553901
perf(applications): check the activations with one call per trial
lenarddome Sep 29, 2026
2980bfd
fix(models): NaN policies and short feedback behave as on main again
lenarddome Sep 29, 2026
3f4894a
bench: re-measure the speed-ups, and update them in the News post
lenarddome Sep 29, 2026
435da17
docs: complete the changelog, and describe the numba fallbacks
lenarddome Sep 29, 2026
313825f
Merge origin/main into feature/numba-integration
lenarddome Sep 29, 2026
d7bf882
ci: use actions/setup-python v7 in the tests, as in the other workflows
lenarddome Sep 29, 2026
2468bae
build: require pandas and matplotlib releases that work with NumPy 2
lenarddome Sep 29, 2026
82bbfd0
fix(applications): reject NaN indices on every platform, not only on x86
lenarddome Sep 29, 2026
68f623b
test: compare to 1e-12 across CPUs, and store the reference without p…
lenarddome Sep 29, 2026
f9e4367
docs(news): results can differ in the last digit on other machines
lenarddome Sep 29, 2026
040c87b
build: require pybads 1.0.5, the first release that works with NumPy 2
lenarddome Sep 29, 2026
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
97 changes: 97 additions & 0 deletions .github/workflows/tests.yml
Original file line number Diff line number Diff line change
@@ -0,0 +1,97 @@
name: tests

# cpm has to give the same results with and without numba, so the test suite
# runs in both setups: with numba installed, and without it.
on:
pull_request:
push:
branches: [main]
workflow_dispatch:

concurrency:
group: tests-${{ github.ref }}
cancel-in-progress: true

permissions:
contents: read

jobs:
test:
name: ${{ matrix.os }}, Python ${{ matrix.python }}, ${{ matrix.numba && 'numba' || 'no numba' }}
runs-on: ${{ matrix.os }}
timeout-minutes: 30
strategy:
fail-fast: false
matrix:
os: [ubuntu-latest]
python: ["3.11", "3.12", "3.13", "3.14"]
numba: [true, false]
include:
- { os: windows-latest, python: "3.14", numba: true }
- { os: windows-latest, python: "3.14", numba: false }
- { os: macos-latest, python: "3.14", numba: true }
- { os: macos-latest, python: "3.14", numba: false }
steps:
- uses: actions/checkout@v7

- uses: actions/setup-python@v7
with:
python-version: ${{ matrix.python }}
cache: pip

- name: Install cpm
run: pip install -e "${{ matrix.numba && '.[numba]' || '.' }}" pytest

- name: Check that numba is used
if: matrix.numba
run: python -c "from cpm.core import _jit; assert _jit.JIT_ENABLED"

- name: Run the tests
run: python -m pytest -q

- name: Run the kernel and application tests with NUMBA_DISABLE_JIT=1
if: matrix.numba
env:
NUMBA_DISABLE_JIT: "1"
run: python -m pytest -q test/models/test_kernels.py test/applications/test_backends.py

lowest:
name: lowest supported versions, Python 3.11
runs-on: ubuntu-latest
timeout-minutes: 30
steps:
- uses: actions/checkout@v7

- uses: actions/setup-python@v7
with:
python-version: "3.11"

- name: Install cpm with the lowest versions pyproject.toml allows
run: |
pip install uv
uv pip install --system --resolution lowest-direct -e ".[numba]" "pytest>=8"

- name: Run the tests
run: python -m pytest -q

prerelease:
# the fast priors in cpm.generators._fast_priors follow scipy's internals,
# so pre-releases show early when a release would change them
name: pre-releases of NumPy, SciPy and pandas
runs-on: ubuntu-latest
timeout-minutes: 30
continue-on-error: true
steps:
- uses: actions/checkout@v7

- uses: actions/setup-python@v7
with:
python-version: "3.14"

- name: Install cpm and the pre-releases
run: |
pip install -e . pytest
pip install --pre --upgrade numpy scipy pandas

- name: Run the tests
run: python -m pytest -q
93 changes: 59 additions & 34 deletions CHANGELOG.md

Large diffs are not rendered by default.

18 changes: 18 additions & 0 deletions CONTRIBUTING.md
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,14 @@ pip install -e .
pip install -e ".[docs]"
```

5. (Optional) Install numba, which compiles the built-in models (`cpm.applications`) and the kernels in `cpm.models.kernels`:

```bash
pip install -e ".[numba]"
```

cpm has to work, with the same results, with and without numba, so changes to the models are tested both ways (see [Testing](#testing)).

## Development workflow (with direct access to repository)

1. Create a branch from `main`.
Expand Down Expand Up @@ -75,6 +83,16 @@ pytest

If you changed specific functionality, also run targeted tests first (for example, under `test/models/` or `test/optimisation/`).

If you changed the models (`cpm.models`, `cpm.applications`) or anything they use, also run the tests with and without numba, and the benchmark suite:

```bash
python scripts/local_tests.py
```

It runs the test suite with numba (if it is installed), with `CPM_DISABLE_JIT=1` (as if numba were not installed), and with `NUMBA_DISABLE_JIT=1`, and then `benchmarks/run.py`.

numba caches compiled code in the `__pycache__` folders under `cpm/`, and only recompiles a function when its own file changes. After changing `cpm/models/kernels.py` (yourself, or by pulling), the compiled loops in `cpm/applications/_sessions.py` still use the old kernels until you delete the cached `*.nbi` and `*.nbc` files; `scripts/local_tests.py` deletes them before it runs.

## Documentation

The documentation is built with [Sphinx](https://www.sphinx-doc.org/) and the [PyData theme](https://pydata-sphinx-theme.readthedocs.io/).
Expand Down
6 changes: 6 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,12 @@ To install the package, run the following command:
pip install cpm-toolbox
```

To fit the built-in models many times faster, also install the optional [numba](https://numba.pydata.org/) dependency, which compiles them; nothing else changes:

```bash
pip install "cpm-toolbox[numba]"
```

Once the package is installed, you can import it in your Python code:

```python
Expand Down
132 changes: 132 additions & 0 deletions benchmarks/cases.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,132 @@
"""
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.

- "python" and "numba": the application, which computes all trials at once, as
plain Python or compiled with numba. Installing numba makes "numba" the default.
- "trial": the application's per-trial `model` function in a plain per-trial
`Wrapper`, which is how users extend an application trial by trial.

With cpm versions before the applications computed all trials at once, only
"trial" is available, and it is the application itself.
"""

import warnings

import numpy as np
import pandas as pd

import cpm
from cpm.generators import Wrapper
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):
return getattr(getattr(cpm.applications, module), name)


def _on_backend(model, backend):
"""The model on `backend`, or None if this version of cpm cannot run it there."""
session = getattr(model, "_session_model", None)
if backend == "trial":
if session is None:
return model # an older cpm: the application is a per-trial Wrapper
return Wrapper(model=model.model, data=model.data, parameters=model.parameters)
if session is None:
return None
session.backend = backend
return model


def build(case, backend="trial"):
"""
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")
data = bandit()
kwargs = dict(dimensions=4, parameters_settings=[[0.3, 0, 1], [4, 0, 10]])
elif case == "HybridMBMF":
cls = _application("reinforcement_learning", "HybridMBMF")
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)
data = risky()
kwargs = dict(parameters_settings=PT_SETTINGS)
else:
raise ValueError(case)
if backend == "numba" and not _jit_enabled():
return None
model = _on_backend(cls(data=data, **kwargs), backend)
if model is None:
return None
observed = data["observed"].to_numpy()
x = np.array([model.parameters[k].value for k in model.parameters.free()], dtype=float)
return model, observed, x


def _jit_enabled():
try:
from cpm.core import _jit
except ImportError:
return False
return _jit.JIT_ENABLED


CASES = ["RLRW", "HybridMBMF", "PTSM", "PTSM1992", "PTSM2025"]
BACKENDS = ["trial", "python", "numba"]
Loading
Loading