diff --git a/CHANGELOG.md b/CHANGELOG.md index 52bda8e..9e03d3e 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -4,6 +4,10 @@ Versions follow [Semantic Versioning](https://semver.org) (`.._metrics-session.csv` written by `process metrics` and `process peri-event --with-metrics`, for go/no-go sessions: `epoch_correlation`, `response_fraction`, `variance_quench`, `signal_fraction` and `reliability_n`, for all trials and per `sdt_type` (`_sdt-`) and go/no-go stimulus (`_stim-go`, `_stim-nogo`). Cue-aligned epochs end at each trial's response (`--no-mask-response` to disable), and lever-aligned epochs run from the cue to the push, with the baseline before the cue (`--cue-baseline`). Options `--min-rt`, `--response-sd` and `--time-warp`. The existing columns are unchanged. The peri-event report's session table shows the new columns with a trial-group selector ([#173](https://github.com/DuguidLab/mesoscopy/issues/173)). + ### Fixed - `process peri-event --event response` wrongly aligned windows to `start_time + response_time`, which falls before the cue since `response_time` is measured from the cue and every trial starts with the ITI. The lever push is now taken at `cue_onset + response_time` ([#176](https://github.com/DuguidLab/mesoscopy/issues/176)). diff --git a/docs/how-to/metrics.md b/docs/how-to/metrics.md index 2829db9..dd365d1 100644 --- a/docs/how-to/metrics.md +++ b/docs/how-to/metrics.md @@ -45,6 +45,9 @@ Pearson correlation between trial traces over the response window) and, for each `_sd` and `_cv` across trials, ignoring trials where the metric is undefined. `onset_n`, `decay_n` and `offset_n` count the trials where each was found. +For go/no-go sessions aligned to the cue, trial start or lever push, the session table also has the +[reliability](#reliability) columns. + Any metric is empty when it is undefined for that trial, for example a decay when the trace never falls back to the fraction, or an onset when the trace never crosses the threshold. @@ -85,6 +88,47 @@ estimate that the `sd` onset threshold depends on. - Session CV is unstable for any metric whose mean sits near zero, which the extrapolated onset often does. Read the SD in that case. +## Reliability + +How consistent a region's response is from trial to trial, over the part of each trial before the lever push, so +that movement does not inflate it. Taken for go/no-go sessions only, and needs the `cue_onset`, `response_time` and +`sdt_type` trials columns (carried over by `process peri-event`, or joined with `--trials`). Without them, or for +reward-aligned windows, the columns are skipped with a warning. + +**Epoch.** The event the windows are aligned to is read from which trials column matches `event_time`. + +- Cue- or trial-start-aligned: the response window, ending at each trial's response. Misses and correct rejections + end at the median response time of the trials with a response. `--no-mask-response` uses the whole response + window. +- Lever-aligned (`--event response`): from the cue to the lever push. The baseline is `--cue-baseline START END` + (default `-1 0`) before each trial's cue, so extract these windows with a `--pre` long enough to reach it (e.g. + `--pre 5`); trials whose baseline falls outside the window have no response fraction or variance quench + baseline, and a warning gives their count. + +Trials with a response time below `--min-rt` (default 0.2 s) are left out. + +**Metrics.** + +| Column | Meaning | +| --- | --- | +| `epoch_correlation` | Mean zero-lag Pearson correlation between every pair of trials, over the samples both have. How alike single trials are. | +| `response_fraction` | Fraction of trials whose mean epoch value exceeds their baseline mean by `--response-sd` (default 2) baseline SDs. | +| `variance_quench` | Across-trial variance over the epoch divided by that over the baseline. Below 1 when variability drops after the event. | +| `signal_fraction` | Fraction of single-trial variance explained by the trial-mean trace. | +| `reliability_n` | Trials used. | + +**Groups.** Each metric is taken over all trials (``), per trial type (`_sdt-hit`, +`_sdt-miss`, `_sdt-false_alarm`, `_sdt-correct_rejection`) and per stimulus (`_stim-go` for hits and +misses, `_stim-nogo` for false alarms and correct rejections). Every column is always written; a group with +fewer than 10 trials is empty. + +**Trial counts.** None of the metrics depends on the number of trials, but all are noisier with fewer, so read +them alongside `reliability_n` when comparing groups. + +**Time warping.** Epochs end at different times. By default, trials are compared sample by sample on the shared +time axis. `--time-warp` instead resamples each epoch onto a common grid from its start to its end, comparing the +shape of the response regardless of response time. + ## As a library ```python @@ -97,4 +141,6 @@ per_trial, per_session = metrics.metrics_tables(perievent, smoothing=5) `metrics.trial_metrics` works on a `(n_trials, n_samples)` array for one region, and the per-metric functions (`peak`, `auc`, `onset_time`, `extrapolated_onset`, `offset_time`, `decay_time`, `trace_correlation`) are -available on their own. +available on their own. `metrics.reliability_metrics` takes the reliability columns for one region, with options +in a `metrics.ReliabilityOptions`; `epoch_correlation`, `response_fraction`, `variance_quench` and +`signal_fraction` work on masked `(n_trials, n_samples)` arrays from `reliability_epochs` and `epoch_traces`. diff --git a/docs/how-to/reports.md b/docs/how-to/reports.md index efc8e0a..7c07b0d 100644 --- a/docs/how-to/reports.md +++ b/docs/how-to/reports.md @@ -33,5 +33,6 @@ panel in the grid. When the `*_metrics.csv` and `*_metrics-session.csv` from `process metrics` sit next to the input, the report adds a metrics section: onset, peak and offset markers on the traces with the across-trial spread of each, per-trial box plots for the selected region, the atlas coloured by a session metric, and the session -table. Peri-event files written without the trials columns have no trial types; pass the trials CSV with +table. For go/no-go sessions the table also shows the reliability metrics, with a selector for the trial group +they are taken over. Peri-event files written without the trials columns have no trial types; pass the trials CSV with `-t/--trials` to recover them. diff --git a/docs/typical-workflow.md b/docs/typical-workflow.md index 0312b12..4c10bfb 100644 --- a/docs/typical-workflow.md +++ b/docs/typical-workflow.md @@ -96,7 +96,9 @@ mesoscopy process metrics /path/to/example-recording_regions_event-cueonset_peri writes `example-recording_regions_event-cueonset_metrics.csv` with the onset time, peak time, amplitude, area under the curve, decay time, offset time and duration of every trial in every region, with the trials columns carried over from the peri-event file, and `example-recording_regions_event-cueonset_metrics-session.csv` with the mean, -SD and CV of each metric across trials per region, and the mean pairwise correlation between trial traces. +SD and CV of each metric across trials per region, and the mean pairwise correlation between trial traces. For +go/no-go sessions the session table also has trial-to-trial reliability metrics over the cue-to-lever epoch, for all +trials, per trial type and per stimulus; see [Response metrics](how-to/metrics.md#reliability). Traces are baseline-subtracted with the mean over `--baseline START END` (default: all pre-event samples), and metrics are taken over `--response START END` (default: all post-event samples). Onset is the first diff --git a/src/mesoscopy/process/__init__.py b/src/mesoscopy/process/__init__.py index 5dcf5f9..9643ab1 100644 --- a/src/mesoscopy/process/__init__.py +++ b/src/mesoscopy/process/__init__.py @@ -23,8 +23,11 @@ from __future__ import annotations +import dataclasses +import functools import os import typing +import warnings from pathlib import Path import click @@ -490,6 +493,74 @@ def _metrics_options(command: Callable[..., None]) -> Callable[..., None]: return command +def _reliability_options(command: Callable[..., None]) -> Callable[..., None]: + """Reliability metric options, passed to the callback as one `reliability` argument. + + Args: + command (Callable[..., None]): Command callback to decorate; takes a `reliability` keyword argument. + + Returns: + Callable[..., None]: The callback with the options attached. + """ + names = [field.name for field in dataclasses.fields(pm.ReliabilityOptions) if field.name != "min_trials"] + + @functools.wraps(command) + def wrapper(*args: typing.Any, **kwargs: typing.Any) -> None: + kwargs["reliability"] = pm.ReliabilityOptions(**{name: kwargs.pop(name) for name in names}) + command(*args, **kwargs) + + options = [ + click.option( + "--mask-response/--no-mask-response", + default=True, + show_default=True, + help="End cue-aligned reliability epochs at each trial's response; trials without one end at the" + " median response time.", + ), + click.option( + "--min-rt", + type=click.FloatRange(min=0), + default=0.2, + show_default=True, + help="Drop trials with a response time below this, in seconds, from the reliability metrics.", + ), + click.option( + "--response-sd", + type=float, + default=2.0, + show_default=True, + help="Baseline SD multiple a trial's mean epoch response must exceed for response_fraction.", + ), + click.option( + "--time-warp", + is_flag=True, + default=False, + help="Resample each reliability epoch onto a common 0-1 grid.", + ), + click.option( + "--cue-baseline", + type=(float, float), + default=(-1.0, 0.0), + show_default=True, + help="Baseline window START END relative to each trial's cue, for lever-aligned reliability metrics.", + ), + ] + for option in reversed(options): + wrapper = option(wrapper) + return wrapper + + +def _echo_warnings(caught: list[warnings.WarningMessage]) -> None: + """Echo the user warnings recorded while computing metrics. + + Args: + caught (list[warnings.WarningMessage]): Warnings recorded by `warnings.catch_warnings(record=True)`. + """ + for warning in caught: + if issubclass(warning.category, UserWarning): + click.echo(f"Warning: {warning.message}") + + def _validate_metrics_options(smoothing: int, extrapolate_range: tuple[float, float]) -> None: """Reject metric options that `trial_metrics` would refuse. @@ -589,6 +660,7 @@ def _write_metrics(trial_table: pd.DataFrame, session_table: pd.DataFrame, out_d help="Also write the response metrics of `process metrics`, joined with TRIALS_PATH. CSV input only.", ) @_metrics_options +@_reliability_options def perievent_cmd( recording_path: str, trials_path: str, @@ -608,6 +680,7 @@ def perievent_cmd( extrapolate_range: tuple[float, float], smoothing: int, decay_fraction: float, + reliability: pm.ReliabilityOptions, ) -> None: """Extract per-trial windows around a behavioural event from a behaviour-aligned recording. @@ -660,7 +733,8 @@ def perievent_cmd( if with_metrics and n_kept == 0: click.echo("No trials kept, skipping metrics.") elif with_metrics: - with timer.Timer(message="Extracting metrics"): + with timer.Timer(message="Extracting metrics"), warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always", UserWarning) tables = pm.metrics_tables( windows, baseline=baseline, @@ -672,7 +746,10 @@ def perievent_cmd( extrapolate_range=extrapolate_range, decay_fraction=decay_fraction, smoothing=smoothing, + event=event, + reliability=reliability, ) + _echo_warnings(caught) _write_metrics(*tables, out_dir, Path(outpath).stem.removesuffix("_perievent")) @@ -839,6 +916,7 @@ def _perievent_csv( help="Trials CSV to join onto the per-trial table by trial_index, for peri-event files without trials columns.", ) @_metrics_options +@_reliability_options def metrics_cmd( path: str, out_dir: str, @@ -852,13 +930,16 @@ def metrics_cmd( extrapolate_range: tuple[float, float], smoothing: int, decay_fraction: float, + reliability: pm.ReliabilityOptions, ) -> None: """Extract per-trial response metrics and per-session variability from peri-event traces. PATH is the long-format *_perievent.csv written by `process peri-event`. Writes _metrics.csv with onset time, peak time, amplitude, area under the curve, decay time, offset time and duration per trial per region, and _metrics-session.csv with the mean, SD and CV of each across trials per region, plus the mean - pairwise trial-trace correlation. + pairwise trial-trace correlation. For go/no-go sessions aligned to the cue, trial start or lever push, the + session table also has trial-to-trial reliability metrics over the cue-to-lever epoch, for all trials, per + sdt_type and per go/no-go stimulus. """ # noqa: DOC501 import pandas as pd @@ -878,7 +959,8 @@ def metrics_cmd( click.echo(f"Loading trials from {trials_path}...") trials = pd.read_csv(trials_path) - with timer.Timer(message="Extracting metrics"): + with timer.Timer(message="Extracting metrics"), warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always", UserWarning) try: tables = pm.metrics_tables( perievent, @@ -892,9 +974,11 @@ def metrics_cmd( extrapolate_range=extrapolate_range, decay_fraction=decay_fraction, smoothing=smoothing, + reliability=reliability, ) except ValueError as error: msg = f"{path}: {error} Expected the columns written by `process peri-event`." raise click.ClickException(msg) from error + _echo_warnings(caught) _write_metrics(*tables, out_dir, Path(path).stem.removesuffix("_perievent")) diff --git a/src/mesoscopy/process/metrics.py b/src/mesoscopy/process/metrics.py index 057a7c2..f6474b1 100644 --- a/src/mesoscopy/process/metrics.py +++ b/src/mesoscopy/process/metrics.py @@ -23,12 +23,15 @@ from __future__ import annotations +import dataclasses import typing import warnings import numpy as np import numpy.typing as npt +from mesoscopy.process import perievent as pev + if typing.TYPE_CHECKING: import pandas as pd @@ -42,6 +45,45 @@ # `extrapolate` fits a line to the rise and takes where it crosses the baseline. ONSET_METHODS = ("sd", "peak", "extrapolate") +# Reliability metrics per trial group, in output column order. +RELIABILITY_METRICS = ( + "epoch_correlation", + "response_fraction", + "variance_quench", + "signal_fraction", + "reliability_n", +) + +# Trial types of a go/no-go session, split by the stimulus shown. +GO_TYPES = ("hit", "miss") +NOGO_TYPES = ("false_alarm", "correct_rejection") + +# Trials columns the reliability metrics need. +RELIABILITY_COLUMNS = frozenset({"cue_onset", "response_time", "sdt_type"}) + + +@dataclasses.dataclass(frozen=True) +class ReliabilityOptions: + """Options for the trial-to-trial reliability metrics. + + Attributes: + mask_response (bool): Mask cue-aligned epochs from each trial's response time. Defaults to True. + min_rt (float): Trials with a response time below this, in seconds, are dropped. Defaults to 0.2. + response_sd (float): Baseline SD multiple a trial's mean epoch response must exceed to count as a + response. Defaults to 2.0. + time_warp (bool): Resample each epoch onto a common 0-1 grid. Defaults to False. + cue_baseline (tuple[float, float]): Baseline window `[start, end)` relative to each trial's cue, for + lever-aligned windows. Defaults to `(-1.0, 0.0)`. + min_trials (int): Groups with fewer trials get NaN. Defaults to 10. + """ + + mask_response: bool = True + min_rt: float = 0.2 + response_sd: float = 2.0 + time_warp: bool = False + cue_baseline: tuple[float, float] = (-1.0, 0.0) + min_trials: int = 10 + def window_mask( time: npt.NDArray[np.float64], start: float, end: float, *, inclusive_end: bool = True @@ -499,9 +541,15 @@ def metrics_tables( extrapolate_range: tuple[float, float] = (0.2, 0.8), decay_fraction: float = 0.5, smoothing: int = 1, + event: str | None = None, + reliability: ReliabilityOptions | None = None, ) -> tuple[pd.DataFrame, pd.DataFrame]: """Per-trial and per-session metric tables from a long-format peri-event table. + For go/no-go sessions, the per-session table also gets the `reliability_metrics` columns. They are skipped + with a warning when the trials columns they need are missing, the session is not go/no-go, or the windows + are reward-aligned. + Args: perievent (pd.DataFrame): Peri-event table with `trial_index`, `event_time`, `time`, `region` and `F` columns, as written by `process peri-event` from a `_regions.csv`. Any other column is taken as @@ -520,11 +568,16 @@ def metrics_tables( extrapolate_range (tuple[float, float], optional): See `trial_metrics`. Defaults to `(0.2, 0.8)`. decay_fraction (float, optional): See `trial_metrics`. Defaults to 0.5. smoothing (int, optional): See `trial_metrics`. Defaults to 1. + event (str | None, optional): Event the windows are aligned to, one of `perievent.EVENTS`. Defaults to + the event matching `event_time`, via `infer_event`. + reliability (ReliabilityOptions | None, optional): Reliability metric options. Defaults to + `ReliabilityOptions()`. Returns: tuple[pd.DataFrame, pd.DataFrame]: The per-trial table, one row per trial per region with `trial_index`, `event_time`, `region`, the `trial_metrics` columns and then the per-trial columns of `perievent`, and the - per-session table, one row per region with `region` and the `session_metrics` columns. + per-session table, one row per region with `region`, the `session_metrics` columns and the + `reliability_metrics` columns. Raises: ValueError: If `perievent` lacks any of `PERIEVENT_COLUMNS`, has no rows, or has a column other than `time`, @@ -556,6 +609,24 @@ def metrics_tables( raise ValueError(msg) events = trial_info[["trial_index", "event_time"]] + info = trial_info + if trials is not None: + info = join_trials(trial_info, trials.drop(columns=extra_columns, errors="ignore")) + reliability = reliability if reliability is not None else ReliabilityOptions() + event = event if event is not None else infer_event(info) + skip = _reliability_skip_reason(info, event) + if skip: + warnings.warn(f"Skipping reliability metrics: {skip}", stacklevel=2) + elif event == "response": + cue = info["cue_onset"].to_numpy(dtype=np.float64) - info["event_time"].to_numpy(dtype=np.float64) + uncovered = int((cue + reliability.cue_baseline[0] < time[0] - _TIME_TOLERANCE_S).sum()) + if uncovered: + warnings.warn( + f"{uncovered} of {len(info)} trials have a pre-cue baseline outside the window; their response" + " fraction and variance quench baselines are NaN. Extract windows with a longer --pre.", + stacklevel=2, + ) + per_trial = [] per_session = [] for region in perievent["region"].unique(): @@ -581,7 +652,10 @@ def metrics_tables( metrics.insert(0, "region", region) per_trial.append(pd.concat([events, metrics], axis=1)) correlation = trace_correlation(traces, time, *response) - per_session.append({"region": region, **session_metrics(metrics, correlation)}) + summary = {"region": region, **session_metrics(metrics, correlation)} + if not skip: + summary |= reliability_metrics(traces, time, info, str(event), baseline, response, reliability) + per_session.append(summary) trial_table = pd.concat(per_trial, ignore_index=True) if extra_columns: @@ -605,3 +679,353 @@ def join_trials(table: pd.DataFrame, trials: pd.DataFrame) -> pd.DataFrame: trials = trials.loc[:, ~trials.columns.str.startswith("Unnamed")].reset_index(drop=True) trials.index.name = "trial_index" return table.merge(trials.reset_index(), on="trial_index", how="left", suffixes=("", "_trial")) + + +# Tolerance when comparing times, in seconds. +_TIME_TOLERANCE_S = 1e-6 + + +def infer_event(trial_info: pd.DataFrame) -> str | None: + """Event the peri-event windows are aligned to, from which trials column `event_time` matches. + + Args: + trial_info (pd.DataFrame): One row per trial with `event_time` and trials columns. + + Returns: + str | None: One of `perievent.EVENTS`, or None when no trials column matches. + """ + event_time = trial_info["event_time"].to_numpy(dtype=np.float64) + for event in pev.EVENTS: + try: + times, index = pev.event_times(trial_info, event) + except KeyError: + continue + if len(index) == len(trial_info) and np.allclose(times, event_time, rtol=0, atol=_TIME_TOLERANCE_S): + return event + return None + + +def _reliability_skip_reason(trial_info: pd.DataFrame, event: str | None) -> str | None: + """Why the reliability metrics cannot be taken. + + Returns: + str | None: The reason, or None when they can be taken. + """ + missing = RELIABILITY_COLUMNS - set(trial_info.columns) + if missing: + return f"no {', '.join(sorted(missing))} column(s); pass the trials CSV." + if "protocol" in trial_info.columns and not trial_info["protocol"].eq("gonogo").all(): + return "only go/no-go sessions are supported." + if event is None: + return "event_time matches no trials column, so the aligned event is unknown." + if event not in {"cue_onset", "trial_start", "response"}: + return f"not defined for {event}-aligned windows." + return None + + +def reliability_epochs( + trial_info: pd.DataFrame, + event: str, + response: tuple[float, float], + mask_response: bool = True, + min_rt: float = 0.2, +) -> tuple[npt.NDArray[np.float64], npt.NDArray[np.float64], npt.NDArray[np.float64], npt.NDArray[np.bool_]]: + """Per-trial reliability epoch `[start, end)`, relative to the event. + + Cue- and trial-start-aligned epochs run over the response window, ending at each trial's response when + `mask_response` is set; trials without a response end at the median response time of those with one. + Lever-aligned (`response`) epochs run from the cue to the lever push. + + Args: + trial_info (pd.DataFrame): One row per trial with `event_time`, `cue_onset` and `response_time` columns. + `response_time` is seconds from the cue. + event (str): `cue_onset`, `trial_start` or `response`. + response (tuple[float, float]): Response window `[start, end]`, in seconds. + mask_response (bool, optional): End cue-aligned epochs at the response. Defaults to True. + min_rt (float, optional): Trials with a response time below this are dropped. Defaults to 0.2. + + Returns: + tuple[npt.NDArray[np.float64], npt.NDArray[np.float64], npt.NDArray[np.float64], npt.NDArray[np.bool_]]: + Epoch start and end, the cue time, all relative to the event, and which trials are kept; each shape + `(n_trials,)`. + + Raises: + ValueError: If `event` is not `cue_onset`, `trial_start` or `response`. + """ + import pandas as pd + + n_trials = len(trial_info) + response_time = pd.to_numeric(trial_info["response_time"], errors="coerce").to_numpy(dtype=np.float64) + responded = response_time >= 0 + keep = ~(responded & (response_time < min_rt)) + cue = trial_info["cue_onset"].to_numpy(dtype=np.float64) - trial_info["event_time"].to_numpy(dtype=np.float64) + + if event == "response": + return cue, np.zeros(n_trials), cue, keep + if event not in {"cue_onset", "trial_start"}: + msg = f"Reliability epochs are not defined for {event!r}; expected cue_onset, trial_start or response." + raise ValueError(msg) + + start = np.full(n_trials, float(response[0])) + # The response window end is inclusive. + window_end = np.nextafter(float(response[1]), np.inf) + if not mask_response: + return start, np.full(n_trials, window_end), cue, keep + timed = responded & keep + median_rt = float(np.median(response_time[timed])) if timed.any() else np.nan + end = np.minimum(cue + np.where(responded, response_time, median_rt), window_end) + return start, end, cue, keep + + +def epoch_traces( + traces: npt.NDArray, time: npt.NDArray[np.float64], start: npt.NDArray[np.float64], end: npt.NDArray[np.float64] +) -> npt.NDArray[np.float64]: + """Traces with every sample outside each trial's `[start, end)` set to NaN. + + Args: + traces (npt.NDArray): Traces of shape `(n_trials, n_samples)`. + time (npt.NDArray[np.float64]): Sample times relative to the event, shape `(n_samples,)`. + start (npt.NDArray[np.float64]): Per-trial epoch start, shape `(n_trials,)`. + end (npt.NDArray[np.float64]): Per-trial epoch end, exclusive, shape `(n_trials,)`. + + Returns: + npt.NDArray[np.float64]: Masked traces, same shape. + """ + inside = (time[None, :] >= start[:, None]) & (time[None, :] < end[:, None]) + return np.where(inside, np.asarray(traces, dtype=np.float64), np.nan) + + +def anchored_window( + traces: npt.NDArray, time: npt.NDArray[np.float64], anchor: npt.NDArray[np.float64], start: float, end: float +) -> npt.NDArray[np.float64]: + """Each trace sampled over `anchor + start <= t < anchor + end` at the recording's sample interval. + + Args: + traces (npt.NDArray): Traces of shape `(n_trials, n_samples)`. + time (npt.NDArray[np.float64]): Sample times relative to the event, shape `(n_samples,)`. + anchor (npt.NDArray[np.float64]): Per-trial anchor time relative to the event, shape `(n_trials,)`. + start (float): Window start relative to the anchor, in seconds. + end (float): Window end relative to the anchor, in seconds. Exclusive. + + Returns: + npt.NDArray[np.float64]: Linearly interpolated samples, shape `(n_trials, n_window)`. Rows are NaN where the + window is not within `time`. + """ + values = np.asarray(traces, dtype=np.float64) + step = float(np.median(np.diff(time))) + offsets = start + step * np.arange(int(np.ceil((end - start) / step - _TIME_TOLERANCE_S))) + window = np.full((values.shape[0], offsets.size), np.nan) + for i in np.flatnonzero(~np.isnan(anchor)): + at = anchor[i] + offsets + if offsets.size and at[0] >= time[0] - _TIME_TOLERANCE_S and at[-1] <= time[-1] + _TIME_TOLERANCE_S: + window[i] = np.interp(at, time, values[i]) + return window + + +def warp_epochs( + epochs: npt.NDArray[np.float64], + time: npt.NDArray[np.float64], + start: npt.NDArray[np.float64], + end: npt.NDArray[np.float64], + n_samples: int, +) -> npt.NDArray[np.float64]: + """Resample each epoch onto `n_samples` points spanning 0 (epoch start) to 1 (epoch end). + + Args: + epochs (npt.NDArray[np.float64]): Masked traces from `epoch_traces`, shape `(n_trials, n_time)`. + time (npt.NDArray[np.float64]): Sample times relative to the event, shape `(n_time,)`. + start (npt.NDArray[np.float64]): Per-trial epoch start, shape `(n_trials,)`. + end (npt.NDArray[np.float64]): Per-trial epoch end, shape `(n_trials,)`. + n_samples (int): Points on the warped grid. + + Returns: + npt.NDArray[np.float64]: Warped epochs, shape `(n_trials, n_samples)`. NaN beyond the samples a trial has, + and for trials with fewer than two samples. + """ + grid = np.linspace(0.0, 1.0, n_samples) + warped = np.full((epochs.shape[0], n_samples), np.nan) + for i in range(epochs.shape[0]): + valid = ~np.isnan(epochs[i]) + if valid.sum() < 2 or not end[i] > start[i]: # noqa: PLR2004 + continue + position = (time[valid] - start[i]) / (end[i] - start[i]) + inside = (grid >= position[0] - _TIME_TOLERANCE_S) & (grid <= position[-1] + _TIME_TOLERANCE_S) + warped[i, inside] = np.interp(grid[inside], position, epochs[i, valid]) + return warped + + +def epoch_correlation(epochs: npt.NDArray[np.float64], min_samples: int = 3) -> float: + """Mean zero-lag Pearson correlation between every pair of trials, over the samples both have. + + Args: + epochs (npt.NDArray[np.float64]): Masked traces of shape `(n_trials, n_samples)`. + min_samples (int, optional): Shared samples a pair needs. Defaults to 3. + + Returns: + float: Mean over pairs with a defined correlation. NaN when there is none. + """ + import pandas as pd + + corr = pd.DataFrame(epochs.T).corr(min_periods=min_samples).to_numpy() + pairs = corr[np.triu_indices_from(corr, k=1)] + return float(np.nanmean(pairs)) if pairs.size and not np.all(np.isnan(pairs)) else float("nan") + + +def response_fraction( + epochs: npt.NDArray[np.float64], baseline: npt.NDArray[np.float64], response_sd: float = 2.0 +) -> float: + """Fraction of trials whose mean epoch value exceeds their baseline mean by `response_sd` baseline SDs. + + Args: + epochs (npt.NDArray[np.float64]): Masked traces of shape `(n_trials, n_samples)`. + baseline (npt.NDArray[np.float64]): Baseline samples of shape `(n_trials, n_baseline)`. + response_sd (float, optional): Baseline SD multiple. Defaults to 2.0. + + Returns: + float: Fraction over trials with a defined epoch mean and baseline. NaN when there is none. + """ + with warnings.catch_warnings(): + warnings.simplefilter("ignore", RuntimeWarning) + response = np.nanmean(epochs, axis=1) - np.nanmean(baseline, axis=1) + threshold = response_sd * np.nanstd(baseline, axis=1, ddof=1) + valid = ~np.isnan(response) & ~np.isnan(threshold) + return float(np.mean(response[valid] > threshold[valid])) if valid.any() else float("nan") + + +def _across_trial_variance(values: npt.NDArray[np.float64]) -> float: + """Mean over samples of the across-trial variance. + + Returns: + float: The mean, over samples with at least two trials. NaN when there are none. + """ + enough = (~np.isnan(values)).sum(axis=0) >= 2 # noqa: PLR2004 + if not enough.any(): + return float("nan") + return float(np.mean(np.nanvar(values[:, enough], axis=0, ddof=1))) + + +def variance_quench(epochs: npt.NDArray[np.float64], baseline: npt.NDArray[np.float64]) -> float: + """Across-trial variance over the epoch divided by across-trial variance over the baseline. + + Args: + epochs (npt.NDArray[np.float64]): Masked traces of shape `(n_trials, n_samples)`. + baseline (npt.NDArray[np.float64]): Baseline samples of shape `(n_trials, n_baseline)`. + + Returns: + float: The ratio; below 1 when variability drops after the event. NaN when either variance is undefined + or the baseline variance is zero. + """ + epoch_variance = _across_trial_variance(epochs) + baseline_variance = _across_trial_variance(baseline) + return epoch_variance / baseline_variance if baseline_variance > 0 else float("nan") + + +def signal_fraction(epochs: npt.NDArray[np.float64]) -> float: + """Fraction of single-trial variance explained by the trial-mean trace. + + Args: + epochs (npt.NDArray[np.float64]): Masked traces of shape `(n_trials, n_samples)`. + + Returns: + float: One minus the residual sum of squares about the trial mean over the total sum of squares, over + samples with at least two trials. NaN when there are none or the traces are constant. + """ + enough = (~np.isnan(epochs)).sum(axis=0) >= 2 # noqa: PLR2004 + values = epochs[:, enough] + if not values.size or np.all(np.isnan(values)): + return float("nan") + residual = np.nansum((values - np.nanmean(values, axis=0)) ** 2) + total = np.nansum((values - np.nanmean(values)) ** 2) + return float(1 - residual / total) if total > 0 else float("nan") + + +def reliability_groups(trial_info: pd.DataFrame) -> dict[str, npt.NDArray[np.bool_]]: + """Trial groups for the reliability metrics, keyed by column suffix. + + Args: + trial_info (pd.DataFrame): One row per trial with an `sdt_type` column. + + Returns: + dict[str, npt.NDArray[np.bool_]]: `""` for all trials, `sdt-` for each of `GO_TYPES` and + `NOGO_TYPES`, then `stim-go` and `stim-nogo`; each a mask of shape `(n_trials,)`. + """ + sdt_type = trial_info["sdt_type"].to_numpy() + groups = {"": np.ones(len(trial_info), dtype=bool)} + for name in (*GO_TYPES, *NOGO_TYPES): + groups[f"sdt-{name}"] = sdt_type == name + groups["stim-go"] = np.isin(sdt_type, GO_TYPES) + groups["stim-nogo"] = np.isin(sdt_type, NOGO_TYPES) + return groups + + +def _group_reliability( + epochs: npt.NDArray[np.float64], baseline: npt.NDArray[np.float64], options: ReliabilityOptions +) -> dict[str, float | int]: + """Reliability metrics for one group. + + Returns: + dict[str, float | int]: One value per `RELIABILITY_METRICS`, NaN below `options.min_trials` trials. + """ + n_trials = epochs.shape[0] + if n_trials < options.min_trials: + return {**{name: float("nan") for name in RELIABILITY_METRICS[:-1]}, "reliability_n": n_trials} + return { + "epoch_correlation": epoch_correlation(epochs), + "response_fraction": response_fraction(epochs, baseline, options.response_sd), + "variance_quench": variance_quench(epochs, baseline), + "signal_fraction": signal_fraction(epochs), + "reliability_n": n_trials, + } + + +def reliability_metrics( + traces: npt.NDArray, + time: npt.NDArray[np.float64], + trial_info: pd.DataFrame, + event: str, + baseline: tuple[float, float], + response: tuple[float, float], + options: ReliabilityOptions | None = None, +) -> dict[str, float | int]: + """Trial-to-trial reliability of one region's response, for all trials and per trial group. + + Metrics are taken over each trial's epoch from `reliability_epochs`: `epoch_correlation`, `response_fraction`, + `variance_quench` and `signal_fraction`, plus `reliability_n`, the trials used. The baseline is `baseline` + relative to the event, or `options.cue_baseline` relative to each trial's cue for lever-aligned windows. Groups + with fewer than `options.min_trials` trials get NaN. + + Args: + traces (npt.NDArray): Traces of shape `(n_trials, n_samples)`, in the row order of `trial_info`. + time (npt.NDArray[np.float64]): Sample times relative to the event, shape `(n_samples,)`. + trial_info (pd.DataFrame): One row per trial with `event_time`, `cue_onset`, `response_time` and + `sdt_type` columns. + event (str): `cue_onset`, `trial_start` or `response`. + baseline (tuple[float, float]): Baseline window `[start, end)` relative to the event, for cue- and + trial-start-aligned windows. + response (tuple[float, float]): Response window `[start, end]`, for cue- and trial-start-aligned windows. + options (ReliabilityOptions | None, optional): Defaults to `ReliabilityOptions()`. + + Returns: + dict[str, float | int]: `` for all trials, then `_` for each `reliability_groups` + group, for each of `RELIABILITY_METRICS`. + + Example: + >>> reliability_metrics(traces, time, trial_info, "cue_onset", baseline=(-1.0, 0.0), response=(0.0, 3.0)) + """ + options = options if options is not None else ReliabilityOptions() + start, end, cue, keep = reliability_epochs(trial_info, event, response, options.mask_response, options.min_rt) + epochs = epoch_traces(traces, time, start, end) + if options.time_warp: + lengths = (~np.isnan(epochs[keep])).sum(axis=1) + lengths = lengths[lengths >= 2] # noqa: PLR2004 + epochs = warp_epochs(epochs, time, start, end, int(np.median(lengths)) if lengths.size else 2) + if event == "response": + base = anchored_window(traces, time, cue, *options.cue_baseline) + else: + base = anchored_window(traces, time, np.zeros(len(cue)), *baseline) + + summary: dict[str, float | int] = {} + for name, mask in reliability_groups(trial_info).items(): + values = _group_reliability(epochs[mask & keep], base[mask & keep], options) + summary |= {(metric if not name else f"{metric}_{name}"): value for metric, value in values.items()} + return summary diff --git a/src/mesoscopy/report/templates/perievent.html b/src/mesoscopy/report/templates/perievent.html index 546eb6f..11b5c2f 100644 --- a/src/mesoscopy/report/templates/perievent.html +++ b/src/mesoscopy/report/templates/perievent.html @@ -153,7 +153,13 @@

Metrics

-
Session metrics (all trials)
+
+ Session metrics (all trials) + + reliability for + + +
@@ -182,6 +188,9 @@

Metrics

const hasMetrics = !!DATA.session_metrics; const MARKERS = {onset_time: {symbol: 'triangle-right', label: 'onset'}, peak_time: {symbol: 'diamond', label: 'peak'}, offset_time: {symbol: 'triangle-left', label: 'offset'}}; const SESSION_COLS = ['region', 'n_trials', 'amplitude_mean', 'amplitude_sd', 'peak_time_mean', 'onset_time_mean', 'offset_time_mean', 'auc_mean', 'duration_mean', 'trace_correlation']; + const RELIABILITY_COLS = ['epoch_correlation', 'response_fraction', 'variance_quench', 'signal_fraction', 'reliability_n']; + const RELIABILITY_GROUPS = [['', 'all trials'], ['sdt-hit', 'hit'], ['sdt-miss', 'miss'], ['sdt-false_alarm', 'false_alarm'], ['sdt-correct_rejection', 'correct_rejection'], ['stim-go', 'go stimulus'], ['stim-nogo', 'no-go stimulus']]; + const hasReliability = hasMetrics && 'epoch_correlation' in DATA.session_metrics[0]; const ATLAS_METRICS = ['amplitude_mean', 'peak_time_mean', 'onset_time_mean', 'offset_time_mean', 'auc_mean', 'duration_mean', 'amplitude_cv', 'trace_correlation']; const state = {types: new Set(TYPES), view: 'mean', split: false, markers: hasMetrics, marks: new Set(['onset_time', 'peak_time']), stat: 'mean', region: DATA.regions[0]}; @@ -388,10 +397,13 @@

Metrics

// ---------- metrics section ---------- function renderSessionTable() { - const cols = SESSION_COLS.filter(c => c in DATA.session_metrics[0]); + // [column, header] pairs; reliability columns follow the selected trial group. + const group = hasReliability ? $('pev-session-group').value : ''; + const cols = SESSION_COLS.filter(c => c in DATA.session_metrics[0]).map(c => [c, c]); + if (hasReliability) cols.push(...RELIABILITY_COLS.map(c => [group ? `${c}_${group}` : c, c])); const fmt = v => typeof v === 'number' ? (Number.isInteger(v) ? v : v.toFixed(3)) : (v ?? ''); - $('pev-session-table').innerHTML = `${cols.map(c => `${c}`).join('')}` + - DATA.session_metrics.map(row => `${cols.map(c => `${fmt(row[c])}`).join('')}`).join('') + ''; + $('pev-session-table').innerHTML = `${cols.map(([, h]) => `${h}`).join('')}` + + DATA.session_metrics.map(row => `${cols.map(([c]) => `${fmt(row[c])}`).join('')}`).join('') + ''; $('pev-session-table').querySelectorAll('tbody tr').forEach(tr => tr.onclick = () => selectRegion(tr.dataset.region)); } function renderMetricAtlas() { @@ -423,6 +435,11 @@

Metrics

buildAtlas($('pev-atlas-metric')); $('pev-metric').innerHTML = ATLAS_METRICS.filter(m => m in DATA.session_metrics[0]).map(m => ``).join(''); $('pev-metric').onchange = renderMetricAtlas; + if (hasReliability) { + $('pev-session-group-wrap').classList.replace('d-none', 'd-inline-flex'); + $('pev-session-group').innerHTML = RELIABILITY_GROUPS.map(([g, label]) => ``).join(''); + $('pev-session-group').onchange = renderSessionTable; + } } // ---------- render ---------- diff --git a/tests/test_process.py b/tests/test_process.py index dd1a0a8..900d6d3 100644 --- a/tests/test_process.py +++ b/tests/test_process.py @@ -2810,3 +2810,373 @@ def test_metrics_cmd_empty_input(metrics_perievent_csv, output_dir, tmp_path): assert str(empty) in result.output assert "no rows" in result.output assert not list(pathlib.Path(output_dir).glob("*_metrics*.csv")) + + +# --------------------------------------------------------------------------- +# reliability metrics +# --------------------------------------------------------------------------- + +GONOGO_TYPES = ["hit"] * 12 + ["miss"] * 6 + ["false_alarm"] * 12 + ["correct_rejection"] * 6 +GONOGO_RT = [0.4, 0.6, 0.8, 1.0, 1.2, 1.4] * 2 + [np.nan] * 6 + [0.5, 0.7, 0.9, 1.1, 1.3, 1.5] * 2 + [np.nan] * 6 + + +def _gonogo_info(event="cue_onset"): + """One row per trial of a go/no-go session with a 5 s ITI; `response_time` is seconds from the cue.""" + n_trials = len(GONOGO_TYPES) + cue = 10.0 * np.arange(n_trials) + 5.0 + response_time = np.asarray(GONOGO_RT) + event_time = {"cue_onset": cue, "trial_start": cue - 5.0, "response": cue + response_time}[event] + keep = ~np.isnan(event_time) + return pd.DataFrame( + { + "trial_index": np.arange(n_trials), + "event_time": event_time, + "start_time": cue - 5.0, + "cue_onset": cue, + "stop_time": np.where(np.isnan(response_time), cue + 2.0, cue + response_time), + "response_time": response_time, + "sdt_type": GONOGO_TYPES, + "protocol": "gonogo", + } + )[keep].reset_index(drop=True) + + +def _gonogo_traces(n_trials, noise=0.1, seed=0): + """A shared response rising after the event, plus independent Gaussian noise per trial.""" + rng = np.random.default_rng(seed) + shape = np.interp(METRICS_GRID, [0.0, 0.5, 1.5], [0.0, 1.0, 0.0]) + return shape + noise * rng.standard_normal((n_trials, len(METRICS_GRID))) + + +def _gonogo_perievent(event="cue_onset"): + """Long-format peri-event table of the go/no-go session for two regions, the second pure noise.""" + info = _gonogo_info(event) + traces = _gonogo_traces(len(info)) + noise = np.random.default_rng(1).standard_normal(traces.shape) + frames = [] + for i, row in info.iterrows(): + for region, values in zip(PERIEVENT_REGIONS, (traces[i], noise[i]), strict=True): + frames.append( + pd.DataFrame( + { + "trial_index": row["trial_index"], + "event_time": row["event_time"], + "time": METRICS_GRID, + "region": region, + "F": values, + } + ) + ) + table = pd.concat(frames, ignore_index=True) + return table.merge(info.drop(columns="event_time"), on="trial_index") + + +class TestInferEvent: + @pytest.mark.parametrize("event", ["cue_onset", "trial_start", "response"]) + def test_matches_event(self, event): + assert pm.infer_event(_gonogo_info(event)) == event + + def test_unknown(self): + info = _gonogo_info() + info["event_time"] += 0.3 + assert pm.infer_event(info) is None + + def test_missing_columns(self): + assert pm.infer_event(_gonogo_info().drop(columns=["cue_onset", "start_time", "stop_time"])) is None + + +class TestReliabilityEpochs: + def test_cue_masked_at_response(self): + start, end, cue, keep = pm.reliability_epochs(_gonogo_info(), "cue_onset", (0.0, 3.0)) + np.testing.assert_allclose(start, 0.0) + np.testing.assert_allclose(cue, 0.0) + np.testing.assert_allclose(end[:12], GONOGO_RT[:12]) + # Trials without a response end at the median of the others. + np.testing.assert_allclose(end[12:18], np.nanmedian(GONOGO_RT)) + assert keep.all() + + def test_trial_start_offsets_by_iti(self): + _, end, cue, _ = pm.reliability_epochs(_gonogo_info("trial_start"), "trial_start", (0.0, 10.0)) + np.testing.assert_allclose(cue, 5.0) + np.testing.assert_allclose(end[:12], 5.0 + np.asarray(GONOGO_RT[:12])) + + def test_capped_at_response_window(self): + _, end, _, _ = pm.reliability_epochs(_gonogo_info(), "cue_onset", (0.0, 0.7)) + assert end.max() == pytest.approx(0.7) + assert end.max() > 0.7 + + def test_unmasked_runs_over_window(self): + _, end, _, _ = pm.reliability_epochs(_gonogo_info(), "cue_onset", (0.0, 3.0), mask_response=False) + assert (end > 3.0).all() + assert (end < 3.0 + 1e-9).all() + + def test_min_rt_drops_fast_trials(self): + _, _, _, keep = pm.reliability_epochs(_gonogo_info(), "cue_onset", (0.0, 3.0), min_rt=0.65) + rt = np.asarray(GONOGO_RT) + np.testing.assert_array_equal(keep, ~(rt < 0.65)) + + def test_lever_runs_from_cue(self): + info = _gonogo_info("response") + start, end, cue, keep = pm.reliability_epochs(info, "response", (0.0, 3.0)) + np.testing.assert_allclose(start, -info["response_time"]) + np.testing.assert_allclose(cue, -info["response_time"]) + np.testing.assert_allclose(end, 0.0) + assert keep.all() + + def test_reward_raises(self): + with pytest.raises(ValueError, match="not defined"): + pm.reliability_epochs(_gonogo_info(), "reward", (0.0, 3.0)) + + +class TestEpochTraces: + def test_masks_outside_epoch(self): + t = np.array([-0.5, 0.0, 0.5, 1.0]) + epochs = pm.epoch_traces(np.ones((2, 4)), t, np.array([0.0, -0.5]), np.array([1.0, np.nan])) + np.testing.assert_array_equal(epochs[0], [np.nan, 1.0, 1.0, np.nan]) + assert np.isnan(epochs[1]).all() + + +class TestAnchoredWindow: + def test_interpolates_relative_to_anchor(self): + t = np.arange(-2.0, 1.01, 0.5) + window = pm.anchored_window(t[None, :].repeat(2, axis=0), t, np.array([0.0, -0.75]), -1.0, 0.0) + np.testing.assert_allclose(window[0], [-1.0, -0.5]) + np.testing.assert_allclose(window[1], [-1.75, -1.25]) + + def test_outside_window_is_nan(self): + t = np.arange(-2.0, 1.01, 0.5) + window = pm.anchored_window(np.ones((2, len(t))), t, np.array([-1.5, np.nan]), -1.0, 0.0) + assert np.isnan(window).all() + + +class TestWarpEpochs: + def test_linear_epochs_share_grid(self): + t = METRICS_GRID + start, end = np.array([0.0, 0.0]), np.array([1.0, 2.0]) + epochs = pm.epoch_traces(np.stack([t, t / 2]), t, start, end) + warped = pm.warp_epochs(epochs, t, start, end, 11) + # Both rise from 0 towards 1 over their own epoch; the last grid point is past the last sample. + np.testing.assert_allclose(warped[:, :10], np.linspace(0, 1, 11)[None, :10].repeat(2, axis=0), atol=1e-9) + assert np.isnan(warped[:, -1]).all() + + def test_too_short_is_nan(self): + t = METRICS_GRID + epochs = pm.epoch_traces(np.ones((1, len(t))), t, np.array([0.0]), np.array([0.01])) + assert np.isnan(pm.warp_epochs(epochs, t, np.array([0.0]), np.array([0.01]), 5)).all() + + +class TestReliabilityMeasures: + def test_identical_trials(self): + epochs = _gonogo_traces(20, noise=0.0) + assert pm.epoch_correlation(epochs) == pytest.approx(1.0) + assert pm.signal_fraction(epochs) == pytest.approx(1.0) + + def test_pure_noise(self): + epochs = np.random.default_rng(0).standard_normal((60, 100)) + assert abs(pm.epoch_correlation(epochs)) < 0.02 + assert pm.signal_fraction(epochs) == pytest.approx(1 / 60, abs=0.01) + + def test_epoch_correlation_uses_shared_samples(self): + a = np.array([1.0, 2.0, 3.0, 4.0, np.nan]) + b = np.array([np.nan, 4.0, 6.0, 8.0, 1.0]) + assert pm.epoch_correlation(np.stack([a, b])) == pytest.approx(1.0) + assert np.isnan(pm.epoch_correlation(np.stack([a, b]), min_samples=4)) + + def test_response_fraction(self): + baseline = np.tile([-0.1, 0.1, -0.1, 0.1], (4, 1)) # SD about 0.115 + epochs = np.array([[1.0, 1.0], [0.1, 0.1], [0.5, np.nan], [np.nan, np.nan]]) + assert pm.response_fraction(epochs, baseline, 2.0) == pytest.approx(2 / 3) + + def test_variance_quench(self): + rng = np.random.default_rng(0) + baseline = rng.standard_normal((200, 25)) + epochs = 0.5 * rng.standard_normal((200, 25)) + assert pm.variance_quench(epochs, baseline) == pytest.approx(0.25, rel=0.1) + assert np.isnan(pm.variance_quench(epochs, np.zeros((200, 25)))) + + +class TestReliabilityGroups: + def test_groups(self): + groups = pm.reliability_groups(_gonogo_info()) + assert list(groups) == [ + "", + "sdt-hit", + "sdt-miss", + "sdt-false_alarm", + "sdt-correct_rejection", + "stim-go", + "stim-nogo", + ] + assert groups[""].sum() == 36 + assert groups["sdt-hit"].sum() == 12 + assert groups["stim-go"].sum() == 18 + assert groups["stim-nogo"].sum() == 18 + + +class TestReliabilityMetrics: + def test_columns_and_min_trials(self): + info = _gonogo_info() + summary = pm.reliability_metrics( + _gonogo_traces(len(info)), METRICS_GRID, info, "cue_onset", (-1.0, 0.0), (0.0, 3.0) + ) + groups = ["", "_sdt-hit", "_sdt-miss", "_sdt-false_alarm", "_sdt-correct_rejection", "_stim-go", "_stim-nogo"] + assert list(summary) == [f"{metric}{group}" for group in groups for metric in pm.RELIABILITY_METRICS] + assert summary["reliability_n"] == 36 + assert summary["reliability_n_sdt-miss"] == 6 + # Six misses are below the default 10-trial minimum. + assert np.isnan(summary["epoch_correlation_sdt-miss"]) + assert summary["epoch_correlation"] > 0.5 + assert summary["epoch_correlation_sdt-hit"] > 0.5 + assert 0 < summary["signal_fraction"] <= 1 + + def test_masking_changes_epochs(self): + info = _gonogo_info() + traces = _gonogo_traces(len(info)) + args = (traces, METRICS_GRID, info, "cue_onset", (-1.0, 0.0), (0.0, 3.0)) + masked = pm.reliability_metrics(*args) + unmasked = pm.reliability_metrics(*args, pm.ReliabilityOptions(mask_response=False)) + assert masked["epoch_correlation"] != unmasked["epoch_correlation"] + # Unmasked over the whole response window, epoch correlation is the trace correlation. + assert unmasked["epoch_correlation"] == pytest.approx(pm.trace_correlation(traces, METRICS_GRID, 0.0, 3.0)) + + def test_min_rt_excludes_trials(self): + info = _gonogo_info() + options = pm.ReliabilityOptions(min_rt=0.65) + summary = pm.reliability_metrics( + _gonogo_traces(len(info)), METRICS_GRID, info, "cue_onset", (-1.0, 0.0), (0.0, 3.0), options + ) + assert summary["reliability_n_sdt-hit"] == 8 + # Two false alarms are also faster than 0.65 s. + assert summary["reliability_n"] == 36 - 6 + + def test_time_warp(self): + info = _gonogo_info() + summary = pm.reliability_metrics( + _gonogo_traces(len(info)), + METRICS_GRID, + info, + "cue_onset", + (-1.0, 0.0), + (0.0, 3.0), + pm.ReliabilityOptions(time_warp=True), + ) + assert np.isfinite(summary["epoch_correlation"]) + + def test_lever_baseline_before_cue(self): + info = _gonogo_info("response") + traces = _gonogo_traces(len(info)) + summary = pm.reliability_metrics( + traces, + METRICS_GRID, + info, + "response", + (-1.0, 0.0), + (0.0, 3.0), + pm.ReliabilityOptions(cue_baseline=(-0.5, 0.0)), + ) + assert summary["reliability_n"] == 24 + assert np.isfinite(summary["variance_quench"]) + # With a baseline reaching before the window, every trial's baseline is NaN. + uncovered = pm.reliability_metrics( + traces, + METRICS_GRID, + info, + "response", + (-1.0, 0.0), + (0.0, 3.0), + pm.ReliabilityOptions(cue_baseline=(-2.0, 0.0)), + ) + assert np.isnan(uncovered["variance_quench"]) + assert np.isnan(uncovered["response_fraction"]) + assert uncovered["epoch_correlation"] == pytest.approx(summary["epoch_correlation"]) + + +class TestMetricsTablesReliability: + def test_gonogo_session(self): + _, per_session = pm.metrics_tables(_gonogo_perievent()) + assert per_session["region"].tolist() == PERIEVENT_REGIONS + assert "epoch_correlation_stim-nogo" in per_session.columns + signal, noise = per_session["epoch_correlation"] + assert signal > 0.5 + assert abs(noise) < 0.1 + + def test_old_columns_unchanged(self): + perievent = _gonogo_perievent() + _, with_reliability = pm.metrics_tables(perievent) + with pytest.warns(UserWarning, match="Skipping reliability"): + _, without = pm.metrics_tables(perievent.drop(columns="sdt_type")) + pd.testing.assert_frame_equal(with_reliability[without.columns], without) + + def test_lever_aligned(self): + _, per_session = pm.metrics_tables(_gonogo_perievent("response"), baseline=(-1.0, 0.0)) + assert per_session["reliability_n"].tolist() == [24, 24] + assert per_session["reliability_n_sdt-miss"].tolist() == [0, 0] + + def test_lever_baseline_outside_window_warns(self): + with pytest.warns(UserWarning, match="pre-cue baseline outside the window"): + pm.metrics_tables(_gonogo_perievent("response"), reliability=pm.ReliabilityOptions(cue_baseline=(-2, 0))) + + def test_trials_supply_columns(self): + perievent = _gonogo_perievent() + trials = _gonogo_info().drop(columns=["trial_index", "event_time"]) + bare = perievent[["trial_index", "event_time", "time", "region", "F"]] + _, joined = pm.metrics_tables(bare, trials=trials) + _, carried = pm.metrics_tables(perievent) + pd.testing.assert_frame_equal(joined, carried) + + @pytest.mark.parametrize( + ("change", "match"), + [ + (lambda table: table.drop(columns="response_time"), "no response_time column"), + (lambda table: table.assign(protocol="2afc"), "only go/no-go"), + (lambda table: table.assign(event_time=table["event_time"] + 0.3), "aligned event is unknown"), + ], + ) + def test_skipped(self, change, match): + with pytest.warns(UserWarning, match=match): + _, per_session = pm.metrics_tables(change(_gonogo_perievent())) + assert "epoch_correlation" not in per_session.columns + + def test_reward_skipped(self): + with pytest.warns(UserWarning, match="not defined for reward"): + _, per_session = pm.metrics_tables(_gonogo_perievent(), event="reward") + assert "epoch_correlation" not in per_session.columns + + +def test_metrics_cmd_reliability(output_dir, tmp_path): + path = tmp_path / "ses-01_regions_event-cueonset_perievent.csv" + _gonogo_perievent().to_csv(path, index=False) + result = CliRunner().invoke( + mesoscopy.cli, + args=f"process metrics {path} -o {output_dir} --no-mask-response --min-rt 0.3 --response-sd 1" + " --time-warp --cue-baseline -0.5 0", + ) + assert result.exit_code == 0, result.output + session = pd.read_csv(pathlib.Path(output_dir) / "ses-01_regions_event-cueonset_metrics-session.csv") + expected = pm.metrics_tables( + pd.read_csv(path), + reliability=pm.ReliabilityOptions( + mask_response=False, + min_rt=0.3, + response_sd=1.0, + time_warp=True, + cue_baseline=(-0.5, 0.0), + ), + )[1] + pd.testing.assert_frame_equal(session, expected, check_dtype=False) + assert session["reliability_n"].tolist() == [36, 36] + + +def test_metrics_cmd_reliability_skipped_warns(metrics_perievent_csv, output_dir): + result = CliRunner().invoke(mesoscopy.cli, args=f"process metrics {metrics_perievent_csv} -o {output_dir}") + assert result.exit_code == 0, result.output + assert "Warning: Skipping reliability metrics: no cue_onset, response_time, sdt_type column(s)" in result.output + + +def test_perievent_cmd_reliability_options_rejected(perievent_regions_csv, perievent_trials_csv, output_dir): + result = CliRunner().invoke( + mesoscopy.cli, + args=f"process peri-event {perievent_regions_csv} {perievent_trials_csv} -o {output_dir} --min-rt -1", + ) + assert result.exit_code == 2 + assert "--min-rt" in result.output diff --git a/tests/test_report.py b/tests/test_report.py index 24fff64..809c384 100644 --- a/tests/test_report.py +++ b/tests/test_report.py @@ -192,6 +192,8 @@ def test_report_cmd_writes_perievent_report(perievent_csv, output_dir): assert payload["trial_index"] == [0, 1, 2, 3, 4] assert payload["event_time"] == [4.5, 9.5, 14.5, 19.5, 24.5] assert len(payload["session_metrics"]) == 3 + assert "epoch_correlation_stim-go" in payload["session_metrics"][0] + assert 'id="pev-session-group"' in html assert {row["region"] for row in payload["trial_metrics"]} == set(REGIONS) assert set(payload["trial_metrics"][0]) == set(pevreport.TRIAL_METRIC_COLUMNS) assert "L_MOp1" in payload["atlas"]["paths"]