Per-spike diagnostics¶
The paper's method: HPD overlap, KL divergence and the rank-based predictive p-value for every spike, thresholds from a baseline period, and flagging.
events
¶
Per-event diagnostics for marked point-process observations.
Spike trains (and other event streams) are usually decoded with a marked
point-process observation model: each event carries a discrete mark (for
spike-sorted data, the identity of the unit that fired) and each mark has a
state-dependent intensity. The functions here turn such a model into the
per-event quantities that the distribution-level diagnostics in
:mod:statespacecheck.state_consistency compare:
- :func:
event_likelihoodnormalizes one event's mark intensity over the state space, giving the single-event likelihood. - :func:
predictive_mark_probabilitiesgives the predictive probability of each mark for the next event. - :func:
event_weighted_predictivegives the state distribution of the next event, weighting the predictive distribution by the total event intensity. - :func:
mark_predictive_pvalueevaluates the predictive check exactly by summing over the finite set of marks. - :func:
event_diagnosticscomputes HPD overlap, KL divergence, and the exact predictive p-value for every event in a recording. - :func:
baseline_thresholdestimates a flagging threshold from a baseline (well-specified) sample of per-event values. - :func:
flag_eventsapplies the paper's flagging rule to every event.
Array conventions follow the rest of the package: state distributions are
(n_events, ...) or (n_time, ...) where ... is one or more spatial
axes, and mark intensity tables are (..., n_marks) over the same spatial
axes.
Classes¶
EventDiagnostics
¶
Bases: NamedTuple
Per-event diagnostic values: HPD overlap, KL divergence and the p-value.
Returned by :func:event_diagnostics and
:func:~statespacecheck.clusterless_event_diagnostics.
Attributes:
| Name | Type | Description |
|---|---|---|
hpd_overlap |
(ndarray, shape(n_events))
|
HPD overlap between the predictive distribution and each event's likelihood. Low values indicate poor local fit. |
kl_divergence |
(ndarray, shape(n_events))
|
KL divergence from the predictive distribution to each event's likelihood. High values indicate poor local fit. |
predictive_pvalue |
(ndarray, shape(n_events))
|
Rank-based predictive p-value of each event's observed mark: exact
from :func: |
likelihood |
np.ndarray, shape (n_events, ...), or None
|
Normalized single-event likelihood of each event over the state space.
|
EventFlags
¶
Bases: NamedTuple
Per-event flags returned by :func:flag_events.
Each field is a boolean array of shape (n_events,) marking events
whose diagnostic is on the misfit side of its threshold, or None if
no threshold was given for that diagnostic.
Attributes:
| Name | Type | Description |
|---|---|---|
hpd_overlap |
np.ndarray of bool, shape (n_events,), or None
|
HPD overlap at or below its threshold. |
kl_divergence |
np.ndarray of bool, shape (n_events,), or None
|
KL divergence at or above its threshold. |
predictive_pvalue |
np.ndarray of bool, shape (n_events,), or None
|
Predictive p-value at or below its cutoff. |
Functions:¶
event_likelihood
¶
Normalize event intensities over the state space.
For a marked point-process observation model, the likelihood contribution
of a single event with mark c is proportional to that mark's intensity
lambda_c(x) (or expected count lambda_c(x) * dt). This function
normalizes each row to sum to 1 over the spatial axes, giving the
single-event likelihood that :func:~statespacecheck.hpd_overlap and
:func:~statespacecheck.kl_divergence compare with the predictive
distribution.
The Poisson exposure term exp(-sum_c lambda_c(x) * dt) is shared by
every event in a time bin, so it is deliberately left out: attaching it to
each event would count it once per event when a bin contains several. A
common bin width dt cancels on normalization, so rates and expected
counts give the same result.
Normalization is done in log space, so rows with tiny but nonzero
intensities keep their shape ([1e-20, 2e-20, 4e-20] becomes
[1/7, 2/7, 4/7]).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
event_intensities
|
(ndarray, shape(n_events, ...))
|
Nonnegative state-dependent intensity (or expected count) of each
event's mark, where |
required |
Returns:
| Name | Type | Description |
|---|---|---|
likelihood |
(ndarray, shape(n_events, ...))
|
Single-event likelihood; each row sums to 1 over the spatial axes. |
Raises:
| Type | Description |
|---|---|
ValueError
|
If the input has no spatial axis, contains negative or non-finite values, or has a row that is zero everywhere (no defined likelihood). |
Examples:
>>> import numpy as np
>>> from statespacecheck import event_likelihood
>>> event_likelihood(np.array([[1.0, 2.0, 1.0]]))
array([[0.25, 0.5 , 0.25]])
Source code in src/statespacecheck/events.py
predictive_mark_probabilities
¶
predictive_mark_probabilities(state_dist: ArrayLike, mark_intensities: ArrayLike) -> DistributionArray
Compute the predictive probability of each mark for the next event.
Mark intensities are averaged over the state distribution and then normalized across marks:
q[c] = sum_x p[x] * lambda_c(x) / sum_d sum_x p[x] * lambda_d(x).
This is the mark distribution of a randomly selected event under the predictive distribution. Normalizing across marks at each state before averaging would drop the weighting by the state-dependent total event rate, and is only equivalent when that total is constant across states.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
state_dist
|
(ndarray, shape(n_events, ...))
|
Predictive state distribution for each event, where |
required |
mark_intensities
|
(ndarray, shape(..., n_marks))
|
Nonnegative intensity (or expected count) of every mark at every state. |
required |
Returns:
| Name | Type | Description |
|---|---|---|
mark_probabilities |
(ndarray, shape(n_events, n_marks))
|
Predictive mark probabilities; each row sums to 1. |
Raises:
| Type | Description |
|---|---|
ValueError
|
If shapes are inconsistent, inputs are negative or non-finite, or a row has zero (or non-finite) total predictive event intensity, for which the mark distribution is undefined. |
Examples:
>>> import numpy as np
>>> from statespacecheck import predictive_mark_probabilities
>>> state = np.array([[0.5, 0.5]])
>>> intensities = np.array([[1.0, 0.0], [1.0, 2.0]]) # (n_bins, n_marks)
>>> predictive_mark_probabilities(state, intensities)
array([[0.5, 0.5]])
Source code in src/statespacecheck/events.py
event_weighted_predictive
¶
Compute the state distribution of the next event.
P_event(x) = Lambda(x) P(x) / sum_u Lambda(u) P(u), where P is the
predictive state distribution and Lambda the ground intensity, the
total event intensity at each state. A randomly chosen event is more
likely to come from states with a higher total event intensity, so the
state of an event is distributed as P weighted by Lambda. With a
constant, positive ground intensity it is the normalized P.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
state_dist
|
(ndarray, shape(n_events, ...))
|
Predictive state distribution for each event, where |
required |
ground_intensity
|
(ndarray, shape(...))
|
Nonnegative total event intensity at every state, over the same
spatial axes. For sorted marks with intensities |
required |
Returns:
| Name | Type | Description |
|---|---|---|
event_weighted |
(ndarray, shape(n_events, ...))
|
Event-weighted state distribution; each row sums to 1. |
Raises:
| Type | Description |
|---|---|
ValueError
|
If shapes are inconsistent, inputs are negative or non-finite, or a row has zero total event intensity under the state distribution, for which the event's state is undefined. |
See Also
predictive_mark_probabilities : The mark distribution of the next event.
Examples:
>>> import numpy as np
>>> from statespacecheck import event_weighted_predictive
>>> state = np.array([[0.5, 0.5]])
>>> event_weighted_predictive(state, np.array([1.0, 3.0]))
array([[0.25, 0.75]])
Source code in src/statespacecheck/events.py
mark_predictive_pvalue
¶
mark_predictive_pvalue(state_dist: ArrayLike, mark_intensities: ArrayLike, observed_marks: ArrayLike) -> DistributionArray
Exact predictive p-value of each event's observed mark.
With a finite set of marks, the predictive check can be evaluated exactly
rather than by Monte Carlo (compare :func:~statespacecheck.predictive_pvalue).
For each event, the p-value is the probability that a mark drawn from the
predictive mark distribution q (:func:predictive_mark_probabilities)
is no more probable than the observed mark:
p = sum_c q[c] * 1{q[c] <= q[observed]}.
Small values mean the observed mark was unexpected given the predictive
state distribution. A large value means only that this mark passes the
check: the least probable mark gets a small p-value on every event even
under a correct model, and marks the prediction makes equally probable all
get p = 1, however unbalanced their observed frequencies. A relative
tolerance on the <= comparison, 16 * eps * n_bins times the
observed mark's probability, absorbs floating-point rounding (each
probability is a sum of n_bins nonnegative terms, accurate to about
n_bins * eps of itself), so marks with equal predictive probability
receive equal p-values across platforms and at any scale, and a mark more
probable than the observed one by more than this tolerance is not counted.
Probabilities stay accurate however small the intensities (an event near the
smallest normal float, about 2.2e-308, is recomputed at a larger scale),
and are resolved to about 1e-30 in absolute terms.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
state_dist
|
(ndarray, shape(n_events, ...))
|
Predictive state distribution for each event, where |
required |
mark_intensities
|
(ndarray, shape(..., n_marks))
|
Nonnegative intensity (or expected count) of every mark at every state. |
required |
observed_marks
|
(ndarray, shape(n_events))
|
Integer index of the observed mark of each event. |
required |
Returns:
| Name | Type | Description |
|---|---|---|
pvalue |
(ndarray, shape(n_events))
|
Exact predictive p-values in |
Raises:
| Type | Description |
|---|---|
ValueError
|
If shapes are inconsistent, marks are out of range, inputs are masked,
or the predictive mark distribution is undefined (see
:func: |
Examples:
>>> import numpy as np
>>> from statespacecheck import mark_predictive_pvalue
>>> state = np.array([[1.0, 0.0], [1.0, 0.0]])
>>> intensities = np.array([[8.0, 1.0, 1.0], [1.0, 1.0, 8.0]]) # (n_bins, n_marks)
>>> mark_predictive_pvalue(state, intensities, np.array([0, 2]))
array([1. , 0.2])
Source code in src/statespacecheck/events.py
616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 641 642 643 644 645 646 647 648 649 650 651 652 653 654 655 656 657 658 659 660 661 662 663 664 665 666 667 668 669 670 671 672 673 674 675 676 677 678 679 680 681 682 683 684 685 686 687 688 689 690 691 692 693 694 695 696 697 698 699 700 701 702 703 704 705 706 707 | |
event_diagnostics
¶
event_diagnostics(predictive: ArrayLike, mark_intensities: ArrayLike, event_time_ind: ArrayLike, event_marks: ArrayLike, *, coverage: float = DEFAULT_COVERAGE, return_likelihood: bool = False, batch_size: int = DEFAULT_EVENT_BATCH_SIZE) -> EventDiagnostics
Compute HPD overlap, KL divergence, and predictive p-value for every event.
Each event is compared with the one-step predictive distribution of its
time bin. The event's likelihood is its mark intensity normalized over the
state space (:func:event_likelihood), and the predictive p-value is the
exact finite-mark check (:func:mark_predictive_pvalue). Events are
processed in batches to bound memory for long recordings.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
predictive
|
(ndarray, shape(n_time, ...))
|
One-step predictive state distribution |
required |
mark_intensities
|
(ndarray, shape(..., n_marks))
|
Nonnegative intensity (or expected count per bin) of every mark at every state, for example each unit's place field. |
required |
event_time_ind
|
(ndarray, shape(n_events))
|
Time-bin index of each event. If several events share a time bin, list each one separately; all are compared with that bin's predictive distribution, so events of the same mark in the same bin receive identical diagnostics. |
required |
event_marks
|
(ndarray, shape(n_events))
|
Mark index of each event (for spike-sorted data, the unit that fired). |
required |
coverage
|
float
|
Coverage probability of the HPD regions. |
0.95
|
return_likelihood
|
bool
|
If True, also return each event's normalized likelihood,
shape |
False
|
batch_size
|
int
|
Number of events processed at once. |
50_000
|
Returns:
| Type | Description |
|---|---|
EventDiagnostics
|
Per-event |
Raises:
| Type | Description |
|---|---|
ValueError
|
If shapes are inconsistent, indices are out of range, or inputs are negative, non-finite or masked; if an event's mark has zero intensity at every position; or if the predictive distribution of an event's time bin puts no mass where any mark has intensity (or the total overflows). |
Examples:
>>> import numpy as np
>>> from statespacecheck import event_diagnostics
>>> predictive = np.array([[0.7, 0.2, 0.1], [0.1, 0.2, 0.7]]) # (n_time, n_bins)
>>> place_fields = np.array([[5.0, 0.1], [1.0, 1.0], [0.1, 5.0]]) # (n_bins, n_marks)
>>> result = event_diagnostics(
... predictive, place_fields, np.array([0, 1]), np.array([0, 0])
... )
>>> result.predictive_pvalue.round(3)
array([1. , 0.172])
Source code in src/statespacecheck/events.py
710 711 712 713 714 715 716 717 718 719 720 721 722 723 724 725 726 727 728 729 730 731 732 733 734 735 736 737 738 739 740 741 742 743 744 745 746 747 748 749 750 751 752 753 754 755 756 757 758 759 760 761 762 763 764 765 766 767 768 769 770 771 772 773 774 775 776 777 778 779 780 781 782 783 784 785 786 787 788 789 790 791 792 793 794 795 796 797 798 799 800 801 802 803 804 805 806 807 808 809 810 811 812 813 814 815 816 817 818 819 820 821 822 823 824 825 826 827 828 829 830 831 832 833 834 835 | |
baseline_threshold
¶
Estimate a flagging threshold from baseline per-event diagnostic values.
Returns the quantile of values pooled from a period (or simulation)
where the model is believed to be well specified. Use a low quantile for
diagnostics where small values indicate misfit (HPD overlap, e.g. 0.01)
and a high quantile where large values do (KL divergence, e.g. 0.99), then
flag events at or beyond the threshold (:func:flag_events). This is the
paper's rule. NaN values are ignored.
Because the comparison is inclusive, every value tied at the threshold is
flagged, and the flagged fraction of the baseline can exceed quantile.
Ties are common: HPD overlap is exactly 1 when one region is nested in the
other and 0 when they are disjoint, and KL divergence is exactly 0 when the
prediction equals the likelihood. If every baseline HPD overlap is 1, the
threshold is 1 and every event at 1 is flagged. Report the flagged fraction of
the baseline with the threshold. It describes the baseline the threshold was
estimated from, not the rate of false alarms to expect elsewhere.
KL divergence is +inf when the prediction and the likelihood have
disjoint support. Such values are allowed: when the requested quantile
falls among them the threshold is +inf, and then only infinite values
are at or above it.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
baseline_values
|
np.ndarray, shape (n_values,) or any shape
|
Baseline diagnostic values; flattened before the quantile is taken. |
required |
quantile
|
float
|
Quantile in |
required |
Returns:
| Name | Type | Description |
|---|---|---|
threshold |
float
|
The requested quantile (linear interpolation) of the baseline values. |
Raises:
| Type | Description |
|---|---|
TypeError
|
If the baseline values are complex. |
ValueError
|
If |
Examples:
>>> import numpy as np
>>> from statespacecheck import baseline_threshold
>>> baseline_threshold(np.arange(101.0), 0.99)
99.0
Source code in src/statespacecheck/events.py
838 839 840 841 842 843 844 845 846 847 848 849 850 851 852 853 854 855 856 857 858 859 860 861 862 863 864 865 866 867 868 869 870 871 872 873 874 875 876 877 878 879 880 881 882 883 884 885 886 887 888 889 890 891 892 893 894 895 896 897 898 899 900 901 902 903 904 905 906 907 908 909 910 911 912 913 914 915 | |
flag_events
¶
flag_events(diagnostics: EventDiagnostics, *, hpd_overlap_threshold: float | None = None, kl_divergence_threshold: float | None = None, pvalue_threshold: float | None = 0.05) -> EventFlags
Flag events whose diagnostics indicate poor local fit.
Applies the paper's rule to each event independently: an event is flagged
when its HPD overlap is at or below hpd_overlap_threshold, its KL
divergence is at or above kl_divergence_threshold, or its
predictive p-value is at or below pvalue_threshold. Each
diagnostic is flagged separately; NaN values are never flagged. Values tied
at a threshold are all flagged (see :func:baseline_threshold).
Thresholds for HPD overlap and KL divergence depend on the model and the
data, so they have no default. In its simulation the paper sets them from a
period where the model is believed to fit, with :func:baseline_threshold
(1st percentile of HPD overlap, 99th percentile of KL divergence); for real
data without such a period it flags HPD overlap at a fixed 0.05 and sets no
KL divergence cutoff. It flags p-values at a fixed 0.05. It recommends HPD overlap and the predictive p-value as the
primary diagnostics and KL divergence as a reference, because KL
divergence is also large for consistent events when the prediction is
broad.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
diagnostics
|
EventDiagnostics
|
Per-event diagnostics from :func: |
required |
hpd_overlap_threshold
|
float
|
Flag HPD overlap at or below this value. Default None (not flagged). |
None
|
kl_divergence_threshold
|
float
|
Flag KL divergence at or above this value. Default None (not flagged). |
None
|
pvalue_threshold
|
float
|
Flag predictive p-values at or below this value. Default 0.05; None skips the p-value. |
0.05
|
Returns:
| Type | Description |
|---|---|
EventFlags
|
Boolean flags for each diagnostic that has a threshold, else None. |
Raises:
| Type | Description |
|---|---|
ValueError
|
If a threshold is NaN. |
Examples:
Thresholds from a baseline period, as in the paper's simulation:
>>> import numpy as np
>>> from statespacecheck import baseline_threshold, event_diagnostics, flag_events
>>> rng = np.random.default_rng(0)
>>> predictive = rng.dirichlet(np.ones(20), size=200) # (n_time, n_bins)
>>> place_fields = rng.gamma(2.0, size=(20, 8)) # (n_bins, n_units)
>>> time_ind, units = rng.integers(0, 200, 500), rng.integers(0, 8, 500)
>>> diagnostics = event_diagnostics(predictive, place_fields, time_ind, units)
>>> baseline = time_ind < 100
>>> flags = flag_events(
... diagnostics,
... hpd_overlap_threshold=baseline_threshold(diagnostics.hpd_overlap[baseline], 0.01),
... kl_divergence_threshold=baseline_threshold(
... diagnostics.kl_divergence[baseline], 0.99
... ),
... )
>>> flags.predictive_pvalue.shape
(500,)
See Also
baseline_threshold : Threshold from a baseline period event_diagnostics : Compute the per-event diagnostics
Source code in src/statespacecheck/events.py
940 941 942 943 944 945 946 947 948 949 950 951 952 953 954 955 956 957 958 959 960 961 962 963 964 965 966 967 968 969 970 971 972 973 974 975 976 977 978 979 980 981 982 983 984 985 986 987 988 989 990 991 992 993 994 995 996 997 998 999 1000 1001 1002 1003 1004 1005 1006 1007 1008 1009 1010 1011 1012 1013 1014 1015 1016 1017 1018 1019 1020 1021 1022 1023 1024 1025 1026 1027 1028 1029 1030 1031 1032 | |