Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
58 commits
Select commit Hold shift + click to select a range
d8183c4
start with some structural builds for the brain explorer module
lenarddome Aug 7, 2024
c5c66e9
transitional commit
lenarddome Aug 19, 2024
29f2f36
start some skeletal preparations
lenarddome Aug 28, 2024
92330b8
Create scavenger.py
lucakosina Jan 18, 2025
6a16c6a
clean: clean and improve documentation structure and class attributes…
lenarddome Jan 23, 2025
eefe548
docs: fix incorrect reference to codebook
lenarddome Jan 23, 2025
0337a66
added documentation
lucakosina Feb 3, 2025
a6bf704
Create spaceObserver.py
lucakosina Feb 28, 2025
9d21133
add group means
lucakosina Mar 11, 2025
dff1382
add group means
lucakosina Mar 11, 2025
64949f6
Update scavenger.py
lucakosina Apr 11, 2025
499a654
updated scavenger
lucakosina Apr 22, 2025
9d562ab
add GH model
lucakosina May 8, 2025
793af25
update GH
lucakosina May 15, 2025
fcfa5dc
add median RTs
lucakosina May 15, 2025
76daa51
Update GoblinHeist_model.py
lucakosina May 20, 2025
9c18787
new SO
lucakosina May 30, 2025
463d45b
add treasurehunt
lucakosina May 30, 2025
ae45a65
add milkyway
lucakosina May 30, 2025
6f5bcf6
update scavenger
lucakosina May 30, 2025
866f5d2
update MW
lucakosina May 30, 2025
9cb5591
Merge branch 'main' into feature/brainexplorer
lenarddome Jun 2, 2025
1707586
Add GoblinHeist model and initial implementation for decision-making …
lenarddome Jun 2, 2025
d766433
docs: add todo list for MBMF
lenarddome Jun 2, 2025
40970db
add exclusion criteria
lucakosina Jun 3, 2025
30000c4
fix data cleaning
lucakosina Jun 3, 2025
3f0fb21
add MBMF model
lucakosina Jun 10, 2025
97e1090
add GH example data
lucakosina Jun 10, 2025
715bffe
adjusted cleaning
lucakosina Jun 10, 2025
10facce
add MFMB data loading
lucakosina Jun 10, 2025
3242d38
fix error
lucakosina Jun 10, 2025
f102c83
update example data
lucakosina Jun 23, 2025
7cd42d7
updated MBMF model
lucakosina Jul 22, 2025
d94a4d1
Update scavenger.py
lucakosina Jul 22, 2025
93baf5a
adapted BE metrics
lucakosina Aug 24, 2025
fdc74d9
new scavenger version
lucakosina Aug 24, 2025
88f7d9f
fix data exclusion
lucakosina Aug 26, 2025
f82882e
updated metrics
lucakosina Aug 27, 2025
b916d7a
Merge branch 'main' into feature/brainexplorer
lenarddome Oct 16, 2025
aecd35e
Merge branch 'main' into feature/brainexplorer
lenarddome Mar 5, 2026
c69f687
clean: organise code contributions after main merge
lenarddome Mar 5, 2026
7f47153
quickfix: update __all__ to include HumbleTeacher in learning module …
lenarddome Mar 6, 2026
478ae1c
fix: correct formatting and spacing in reinforcement learning classes
lenarddome Mar 12, 2026
d4a7830
feat: add perceptual decision making module and update imports
lenarddome Mar 12, 2026
56380eb
feat: add tests for SpaceObserver functionality and data handling
lenarddome Mar 18, 2026
252f415
chore: update .gitignore to include .vscode and tmp directories
lenarddome Mar 18, 2026
36486e5
fix: improve time_of_day logic in SpaceObserver and enhance related t…
lenarddome Mar 19, 2026
35ca5c3
fix: update file reading logic, fix time of day classification logic,…
lenarddome Mar 19, 2026
957850c
Merge branch 'main' into feature/brainexplorer
lenarddome Oct 5, 2026
f4bb6b6
fix(brainexplorer): fix metrics, exclusions and codebooks
lenarddome Oct 5, 2026
69c2ef1
test(brainexplorer): add unit tests for every task
lenarddome Oct 5, 2026
73bb116
docs(api): add the brainexplorer reference page
lenarddome Oct 5, 2026
d64f86f
fix(optimisation): infer ppt_identifier from the grouping of the data
lenarddome Oct 5, 2026
032b22b
refactor(brainexplorer): read CSV files only, like the data in cpm
lenarddome Oct 5, 2026
40134c5
fix(optimisation): explain a ppt_identifier that is not a column
lenarddome Oct 5, 2026
2889415
fix(generators): restore the log prior of values outside the support
lenarddome Oct 5, 2026
1166277
docs(datasets): document load_model_based_model_free
lenarddome Oct 5, 2026
47cc34b
chore: bump the version to 0.26.0.dev1
lenarddome Oct 5, 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
6 changes: 5 additions & 1 deletion CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@ All notable changes to this project will be documented in this file.
The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.1.0/),
and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).

## [0.26.0.dev0] - 2026-09-28
## [0.26.0.dev1] - 2026-10-05

### Added

Expand All @@ -18,6 +18,9 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
- Added `cpm.models.learning.SARSATrace`, a SARSA learning rule with an eligibility trace
- Added test units for `HybridMBMF` and `SARSATrace`
- Added a two-step task example replicating Smid et al. (2022), with its data in `cpm.datasets.load_two_step_data`
- Added `cpm.brainexplorer`, which computes descriptive statistics and applies the exclusion criteria for the BrainExplorer games Space Observer, Scavenger, Treasure Hunt, and Milky Way and Pirate Market, with a page in the API reference
- Added test units for `cpm.brainexplorer`
- Added `cpm.datasets.load_model_based_model_free`, two-step task data from the BrainExplorer game Goblin Heist
- Added an example that fits causal ratings in a blocking experiment, recreating Figures 1 and 2 of Spicer et al. (2021), with its data in `cpm.datasets.load_blocking_data`; based on an earlier version by @chotong
- Added `cpm.generators.SessionWrapper`, a `Wrapper` for models that compute all trials of a participant at once
- Added `cpm.models.kernels`, the formulas of the `cpm.models` classes as functions that numba can compile, used by the classes and the built-in applications
Expand Down Expand Up @@ -49,6 +52,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0

### Fixed

- Fixed the optimisers and `cpm.applications.signal_detection.EstimatorMetaD` raising `UnboundLocalError` for a pandas DataFrame without `ppt_identifier`; they now raise a `ValueError` that explains what to pass, or a `KeyError` that lists the columns if `ppt_identifier` is not one of them, and take `ppt_identifier` from the grouping of data grouped by a single column, such as `data.groupby("ppt")`, so that the fits record the participants
- Fixed `cpm.hierarchical.EmpiricalBayes` and `cpm.hierarchical.VariationalBayes` discarding the estimated population priors after the first evaluation of each fit; hierarchical results change
- Fixed the random starting priors of later chains of `EmpiricalBayes` and `VariationalBayes` being length-one arrays
- Fixed `numpy.random.seed` not reproducing the later chains of `EmpiricalBayes` and `VariationalBayes`, whose starting priors now also respect the lower bound and work with infinite bounds
Expand Down
1 change: 1 addition & 0 deletions cpm/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,5 +21,6 @@
from . import hierarchical
from . import utils
from . import datasets
from . import brainexplorer

del os, sys, warnings, logging
2 changes: 1 addition & 1 deletion cpm/__version__.py
Original file line number Diff line number Diff line change
@@ -1 +1 @@
__version__ = "0.26.0.dev0"
__version__ = "0.26.0.dev1"
6 changes: 3 additions & 3 deletions cpm/applications/signal_detection.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,7 @@
from scipy.stats import norm, multivariate_normal
from scipy.special import ndtr

from cpm.core.optimisers import numerical_hessian, prepare_data
from cpm.core.optimisers import group_identifier, numerical_hessian, prepare_data
from cpm.core.data import detailed_pandas_compiler, decompose
from cpm.core.parallel import detect_cores, execute_parallel
from cpm.utils.metad import count_trials, bin_ratings
Expand Down Expand Up @@ -428,7 +428,7 @@ class EstimatorMetaD:
* 3 : display progress during iterations (more complete report).

ppt_identifier : str, optional
Identifier for participants in the data. If None, the default identifier will be used.
The column that identifies the participants in the data. Required if `data` is a pandas.DataFrame. If `data` is grouped by a single column, such as `data.groupby("ppt")`, it defaults to the name of that column.
ignore_invalid : bool, default False
If True, invalid confidence ratings will be ignored during binning. If False, an error will be raised if invalid ratings are found. We recommend setting this to False (Default).
**kwargs : additional keyword arguments
Expand Down Expand Up @@ -464,7 +464,7 @@ def __init__(
):
self.data = data
self.bins = bins
self.ppt_identifier = ppt_identifier
self.ppt_identifier = group_identifier(data, ppt_identifier)
self.data, self.participants, self.groups, self.__pandas__ = prepare_data(
data, self.ppt_identifier
)
Expand Down
4 changes: 4 additions & 0 deletions cpm/brainexplorer/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,4 @@
from cpm.brainexplorer import bandits
from cpm.brainexplorer import perceptual_decision_making
from cpm.brainexplorer import risky_decision_making
from cpm.brainexplorer import information_gathering
70 changes: 70 additions & 0 deletions cpm/brainexplorer/_utils.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,70 @@
"""Helpers shared by the BrainExplorer task classes."""

import os
import warnings
from contextlib import contextmanager

import numpy as np
import pandas as pd


def load_data(source, name, **kwargs):
"""
Read the data of a task from a CSV file, or copy a DataFrame.

`source` is a pandas DataFrame, or the path to a CSV file. Keyword arguments
are passed on to :func:`pandas.read_csv`.
"""
if source is None:
raise ValueError(f"{name} needs data: a pandas DataFrame or the path to a CSV file.")
if isinstance(source, pd.DataFrame):
return source.copy()
return pd.read_csv(os.fspath(source), header=0, **kwargs)


def require(data, columns, name):
"""Raise a `KeyError` that names every required column missing from `data`."""
missing = [column for column in columns if column not in data.columns]
if missing:
raise KeyError(f"{name} needs the columns {missing}, which are missing from the data.")


def time_of_day(hour):
"""Name the time of day of an hour: night (0-6h), morning (6-12h), afternoon (12-18h) or evening (18-24h)."""
if hour < 6:
return "night"
if hour < 12:
return "morning"
if hour < 18:
return "afternoon"
return "evening"


def session_time(date):
"""The date, day of the week, clock time and time of day of the first trial of a session."""
date = pd.Timestamp(pd.to_datetime(date, format="ISO8601")) if isinstance(date, str) else pd.Timestamp(date)
return {
"date": date,
"day_of_week": date.day_name(),
"time": date.time(),
"time_of_day": time_of_day(date.hour),
}


@contextmanager
def quiet_empty_slices():
"""Silence the warnings numpy raises for the mean, median or SD of an empty selection, which are NaN by design."""
with warnings.catch_warnings():
warnings.filterwarnings("ignore", message="Mean of empty slice")
warnings.filterwarnings("ignore", message="All-NaN slice encountered")
warnings.filterwarnings("ignore", message="Degrees of freedom <= 0")
warnings.filterwarnings("ignore", message="invalid value encountered")
yield


def proportion(condition, mask):
"""The proportion of trials in `mask` for which `condition` is true, or NaN if `mask` selects no trials."""
mask = np.asarray(mask, dtype=bool)
if not mask.any():
return np.nan
return float(np.mean(np.asarray(condition, dtype=bool)[mask]))
243 changes: 243 additions & 0 deletions cpm/brainexplorer/bandits.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,243 @@
import numpy as np
import pandas as pd

from ._utils import load_data, quiet_empty_slices, require, session_time

__all__ = ["MilkyWay"]

COLUMNS = ["userID", "trial_type", "run", "date", "correct", "outchosen", "obt_min_forg", "rep", "WSLS_v1", "WSLS_v2"]

## the two games, by their trial_type in the data: (name, suffix of the metrics)
GAMES = {"reward": ("Milky Way", "MW"), "punish": ("Pirate Market", "PM")}

## the metrics that are compared between the games: (name of the difference, metric without its suffix)
DIFFERENCES = [
("accuracy_diff", "accuracy"),
("outcome_diff", "mean_outcome"),
("reward_diff_obt_forg_diff", "reward_diff_obt_forg"),
("prop_same_choice_diff", "prop_same_choice"),
("prop_WSLS_1_diff", "prop_WSLS_1"),
("prop_WSLS_2_diff", "prop_WSLS_2"),
]


class MilkyWay:
"""
Compute descriptive statistics from the two-armed bandit tasks in BrainExplorer, *Milky Way* (reward) and *Pirate Market* (punishment).

The metrics of the two games are computed separately, and :meth:`difference_metrics` compares them within participants.

Parameters
----------
filepath : str, os.PathLike or pandas.DataFrame
The data, as a DataFrame or the path to a CSV file. The column names must follow the convention in Notes.

Attributes
----------
MW_data, PM_data : pandas.DataFrame
The Milky Way and Pirate Market trials that pass the trial-level exclusion criteria.
results_MW, results_PM : pandas.DataFrame
The metrics of each participant in each game, filled by :meth:`metrics`.
results_diff : pandas.DataFrame
The differences between the games, filled by :meth:`difference_metrics`.
cleanedresults_MW, cleanedresults_PM, cleanedresults_diff : pandas.DataFrame
The results of the participants that pass the participant-level exclusion criteria, filled by :meth:`clean_data`.
deleted_participants_MW, deleted_participants_PM, deleted_participants_diff : int
The number of participants that :meth:`clean_data` excluded from each table.
codebook : dict
The description of each column of the metrics.

Examples
--------
>>> from cpm.brainexplorer.bandits import MilkyWay
>>> milky_way = MilkyWay("2025-02-20_MilkyWay_Data.csv")
>>> results_MW, results_PM = milky_way.metrics()
>>> differences = milky_way.difference_metrics()
>>> cleaned_MW, cleaned_PM, cleaned_diff = milky_way.clean_data()

Notes
-----
The data must contain the following columns:

- ``userID``: the unique identifier of the participant.
- ``trial_type``: the game, ``"reward"`` for Milky Way or ``"punish"`` for Pirate Market.
- ``run``: the attempt number of the participant.
- ``date``: the date and time of the trial.
- ``correct``: whether the choice was correct (1) or incorrect (0).
- ``outchosen``: the outcome of the chosen option.
- ``obt_min_forg``: the obtained minus the foregone outcome.
- ``rep``: whether the choice repeated the previous one (1) or not (0).
- ``WSLS_v1``, ``WSLS_v2``: whether the choice was a win-stay/lose-shift choice, in two versions of the definition.

Only the first attempt of each participant in each game is kept, which is not necessarily ``run`` 1 for Pirate Market.

Response times are not analysed. They would have to be corrected for the average time between the previous choice and the presentation of the new stimuli, 4200 ms, which can make them negative.
"""

def __init__(self, filepath=None):
data = load_data(filepath, "MilkyWay")
require(data, COLUMNS, "MilkyWay")
self.data = data
self.MW_data = _first_attempt(data[data["trial_type"] == "reward"])
self.PM_data = _first_attempt(data[data["trial_type"] == "punish"])

self.results_MW = pd.DataFrame()
self.results_PM = pd.DataFrame()
self.results_diff = pd.DataFrame()

self.codebook = {
"userID": "Unique identifier for each participant",
"n_trials": "Number of trials completed by the participant",
"trial_type": "Type of trial (Milky Way or Pirate Market)",
"date": "Date and time of the first trial",
"day_of_week": "Day of the week of the first trial",
"time": "Clock time of the first trial",
"time_of_day": "Time of day of the first trial (night: 0-6h, morning: 6-12h, afternoon: 12-18h, evening: 18-24h)",
}
for name, suffix in GAMES.values():
self.codebook.update(
{
f"accuracy_{suffix}": f"Mean accuracy for {name} trials",
f"mean_outcome_{suffix}": f"Mean outcome for {name} trials",
f"reward_diff_obt_forg_{suffix}": f"Mean difference between obtained and foregone reward for {name} trials",
f"prop_same_choice_{suffix}": f"Proportion of same choice in {name} trials",
f"prop_WSLS_1_{suffix}": f"Proportion of win-stay/lose-shift choices (version 1) for {name} trials",
f"prop_WSLS_2_{suffix}": f"Proportion of win-stay/lose-shift choices (version 2) for {name} trials",
}
)
self.codebook.update(
{
"accuracy_diff": "Difference in accuracy: Milky Way minus Pirate Market",
"outcome_diff": "Difference in mean outcome: Milky Way minus Pirate Market",
"reward_diff_obt_forg_diff": "Difference in the mean obtained minus foregone reward: Milky Way minus Pirate Market",
"prop_same_choice_diff": "Difference in the proportion of same choices: Milky Way minus Pirate Market",
"prop_WSLS_1_diff": "Difference in the proportion of win-stay/lose-shift choices (version 1): Milky Way minus Pirate Market",
"prop_WSLS_2_diff": "Difference in the proportion of win-stay/lose-shift choices (version 2): Milky Way minus Pirate Market",
}
)

def metrics(self):
"""
Compute the metrics of each participant in each game.

Returns
-------
tuple of pandas.DataFrame
The metrics of Milky Way and of Pirate Market, one row per participant, also stored in `results_MW` and `results_PM`. The columns are described in `codebook`.
"""
self.results_MW = _game_metrics(self.MW_data, *GAMES["reward"])
self.results_PM = _game_metrics(self.PM_data, *GAMES["punish"])
return self.results_MW, self.results_PM

def difference_metrics(self):
"""
Compute the difference between the metrics of Milky Way and Pirate Market, for each participant who played both.

Runs :meth:`metrics` first if it has not been run yet.

Returns
-------
pandas.DataFrame
The differences, Milky Way minus Pirate Market, one row per participant, also stored in `results_diff`. The columns are described in `codebook`.
"""
if self.results_MW.empty and self.results_PM.empty:
self.metrics()
columns = ["userID"] + [name for name, _ in DIFFERENCES]
if self.results_MW.empty or self.results_PM.empty:
self.results_diff = pd.DataFrame(columns=columns)
return self.results_diff
both = self.results_MW.merge(self.results_PM, on="userID", sort=False)
for name, metric in DIFFERENCES:
both[name] = both[f"{metric}_MW"] - both[f"{metric}_PM"]
self.results_diff = both[columns].reset_index(drop=True)
return self.results_diff

def clean_data(self):
"""
Exclude participants with the participant-level exclusion criteria.

Runs :meth:`metrics` and :meth:`difference_metrics` first if they have not been run yet.

Returns
-------
tuple of pandas.DataFrame
The metrics of Milky Way, of Pirate Market, and their differences, for the participants that pass the criteria, also stored in `cleanedresults_MW`, `cleanedresults_PM` and `cleanedresults_diff`.

Notes
-----
A participant is excluded from the metrics of a game if

- they made the same choice on at least 95% of the trials,
- their accuracy is missing,
- they played more than 72 trials of the game (due to a technical error).

The differences are kept only for participants who pass the criteria in both games.
"""
if self.results_MW.empty and self.results_PM.empty:
self.metrics()
if self.results_diff.empty:
self.difference_metrics()

self.cleanedresults_MW = _clean_game(self.results_MW, self.MW_data, "MW")
self.cleanedresults_PM = _clean_game(self.results_PM, self.PM_data, "PM")
self.deleted_participants_MW = _n_users(self.results_MW) - _n_users(self.cleanedresults_MW)
self.deleted_participants_PM = _n_users(self.results_PM) - _n_users(self.cleanedresults_PM)

both = set(_users(self.cleanedresults_MW)) & set(_users(self.cleanedresults_PM))
self.cleanedresults_diff = self.results_diff[self.results_diff["userID"].isin(both)].copy()
self.deleted_participants_diff = _n_users(self.results_diff) - _n_users(self.cleanedresults_diff)

return self.cleanedresults_MW, self.cleanedresults_PM, self.cleanedresults_diff

def get_codebook(self):
"""
Return the codebook, which describes each column of the metrics.

Returns
-------
dict
The description of each column, keyed by column name.
"""
return self.codebook


def _first_attempt(data):
"""The trials of the first attempt of each participant, which is their lowest `run`."""
first = data.groupby("userID")["run"].transform("min")
return data[data["run"] == first].reset_index(drop=True)


def _game_metrics(data, name, suffix):
"""The metrics of each participant in one game."""
rows = []
with quiet_empty_slices():
for user_id, user_data in data.groupby("userID"):
row = {"userID": user_id, "n_trials": len(user_data), "trial_type": name}
row.update(session_time(user_data["date"].iloc[0]))
row[f"accuracy_{suffix}"] = np.nanmean(user_data["correct"])
row[f"mean_outcome_{suffix}"] = np.nanmean(user_data["outchosen"])
row[f"reward_diff_obt_forg_{suffix}"] = np.nanmean(user_data["obt_min_forg"])
row[f"prop_same_choice_{suffix}"] = np.nanmean(user_data["rep"])
row[f"prop_WSLS_1_{suffix}"] = np.nanmean(user_data["WSLS_v1"])
row[f"prop_WSLS_2_{suffix}"] = np.nanmean(user_data["WSLS_v2"])
rows.append(row)
return pd.DataFrame(rows)


def _clean_game(results, data, suffix):
"""The metrics of one game, for the participants that pass the exclusion criteria."""
if results.empty:
return results.copy()
trials = data.groupby("userID").size()
keep = results["userID"].isin(trials.index[trials <= 72])
keep &= results[f"prop_same_choice_{suffix}"] < 0.95
keep &= results[f"accuracy_{suffix}"].notna()
return results[keep].copy()


def _users(results):
return results["userID"].unique() if "userID" in results else []


def _n_users(results):
return len(_users(results))
Loading
Loading