nltools 0.6.0.dev0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (95) hide show
  1. nltools/__init__.py +55 -0
  2. nltools/algorithms/__init__.py +90 -0
  3. nltools/algorithms/alignment/__init__.py +21 -0
  4. nltools/algorithms/alignment/procrustes.py +565 -0
  5. nltools/algorithms/alignment/srm.py +758 -0
  6. nltools/algorithms/backends.py +1059 -0
  7. nltools/algorithms/corrections.py +177 -0
  8. nltools/algorithms/decoding.py +327 -0
  9. nltools/algorithms/inference/__init__.py +50 -0
  10. nltools/algorithms/inference/bootstrap.py +1386 -0
  11. nltools/algorithms/inference/correlation.py +373 -0
  12. nltools/algorithms/inference/intersubject.py +422 -0
  13. nltools/algorithms/inference/isc.py +1554 -0
  14. nltools/algorithms/inference/matrix.py +602 -0
  15. nltools/algorithms/inference/one_sample.py +288 -0
  16. nltools/algorithms/inference/random.py +122 -0
  17. nltools/algorithms/inference/timeseries.py +347 -0
  18. nltools/algorithms/inference/two_sample.py +212 -0
  19. nltools/algorithms/inference/utils.py +58 -0
  20. nltools/algorithms/inference/validation.py +282 -0
  21. nltools/algorithms/neighborhoods.py +207 -0
  22. nltools/algorithms/outliers.py +308 -0
  23. nltools/algorithms/regression.py +83 -0
  24. nltools/algorithms/signal.py +303 -0
  25. nltools/algorithms/similarity.py +234 -0
  26. nltools/algorithms/validation.py +151 -0
  27. nltools/cross_validation.py +72 -0
  28. nltools/data/__init__.py +30 -0
  29. nltools/data/adjacency/__init__.py +875 -0
  30. nltools/data/adjacency/io.py +111 -0
  31. nltools/data/adjacency/modeling.py +569 -0
  32. nltools/data/adjacency/plotting.py +174 -0
  33. nltools/data/adjacency/state.py +349 -0
  34. nltools/data/adjacency/stats.py +596 -0
  35. nltools/data/adjacency/utils.py +79 -0
  36. nltools/data/atlases/__init__.py +23 -0
  37. nltools/data/atlases/labeling.py +158 -0
  38. nltools/data/atlases/loading.py +76 -0
  39. nltools/data/atlases/registry.py +96 -0
  40. nltools/data/atlases/reporting.py +456 -0
  41. nltools/data/braindata/__init__.py +2170 -0
  42. nltools/data/braindata/analysis.py +1381 -0
  43. nltools/data/braindata/bootstrap.py +398 -0
  44. nltools/data/braindata/io.py +896 -0
  45. nltools/data/braindata/modeling.py +594 -0
  46. nltools/data/braindata/plotting.py +501 -0
  47. nltools/data/braindata/prediction.py +1250 -0
  48. nltools/data/braindata/utils.py +348 -0
  49. nltools/data/braindata/validation.py +197 -0
  50. nltools/data/braindata/viewer.js +266 -0
  51. nltools/data/braindata/viewer.py +770 -0
  52. nltools/data/combine.py +27 -0
  53. nltools/data/designmatrix/__init__.py +1032 -0
  54. nltools/data/designmatrix/append.py +518 -0
  55. nltools/data/designmatrix/diagnostics.py +248 -0
  56. nltools/data/designmatrix/io.py +356 -0
  57. nltools/data/designmatrix/plotting.py +291 -0
  58. nltools/data/designmatrix/regressors.py +463 -0
  59. nltools/data/designmatrix/transforms.py +200 -0
  60. nltools/data/designmatrix/utils.py +350 -0
  61. nltools/data/ownership.py +129 -0
  62. nltools/data/results.py +291 -0
  63. nltools/data/roc/__init__.py +398 -0
  64. nltools/data/simulator/__init__.py +927 -0
  65. nltools/data/simulator/haxby.py +124 -0
  66. nltools/data/validation.py +83 -0
  67. nltools/datasets.py +218 -0
  68. nltools/io/__init__.py +10 -0
  69. nltools/io/events.py +67 -0
  70. nltools/io/h5.py +246 -0
  71. nltools/mask.py +403 -0
  72. nltools/models/__init__.py +11 -0
  73. nltools/models/glm.py +543 -0
  74. nltools/models/results.py +49 -0
  75. nltools/models/ridge.py +1303 -0
  76. nltools/models/validation.py +26 -0
  77. nltools/plotting/__init__.py +32 -0
  78. nltools/plotting/adjacency.py +421 -0
  79. nltools/plotting/brain.py +669 -0
  80. nltools/plotting/decomposition.py +111 -0
  81. nltools/plotting/prediction.py +110 -0
  82. nltools/resources/covariates_example.csv +161 -0
  83. nltools/resources/onsets_example.csv +40 -0
  84. nltools/templates/__init__.py +51 -0
  85. nltools/templates/config.py +144 -0
  86. nltools/templates/fetch.py +260 -0
  87. nltools/templates/matching.py +183 -0
  88. nltools/templates/paths.py +106 -0
  89. nltools/templates/registry.py +25 -0
  90. nltools/utils.py +230 -0
  91. nltools/version.py +13 -0
  92. nltools-0.6.0.dev0.dist-info/METADATA +95 -0
  93. nltools-0.6.0.dev0.dist-info/RECORD +95 -0
  94. nltools-0.6.0.dev0.dist-info/WHEEL +4 -0
  95. nltools-0.6.0.dev0.dist-info/licenses/LICENSE +21 -0
@@ -0,0 +1,308 @@
1
+ """Outlier detection, robust statistics, and data normalization."""
2
+
3
+ import numpy as np
4
+ import polars as pl
5
+ import nibabel as nib
6
+
7
+
8
+ def zscore(data):
9
+ """Z-score every column of a Polars or pandas DataFrame/Series.
10
+
11
+ Pandas inputs are converted to Polars; the result is always Polars (a
12
+ DataFrame for DataFrame input, a Series for Series input).
13
+
14
+ Args:
15
+ data (pl.DataFrame | pl.Series | pd.DataFrame | pd.Series): Data to z-score.
16
+
17
+ Returns:
18
+ pl.DataFrame | pl.Series: Same shape as the input, each column z-scored
19
+ with the sample standard deviation (ddof=1).
20
+ """
21
+ import pandas as pd
22
+
23
+ return_series = False
24
+ if isinstance(data, pd.DataFrame):
25
+ df = pl.from_pandas(data)
26
+ elif isinstance(data, pd.Series):
27
+ df = pl.DataFrame({data.name or "0": data})
28
+ return_series = True
29
+ elif isinstance(data, pl.DataFrame):
30
+ df = data
31
+ elif isinstance(data, pl.Series):
32
+ df = pl.DataFrame({data.name or "0": data})
33
+ return_series = True
34
+ else:
35
+ raise ValueError("Data must be a Polars or pandas DataFrame or Series")
36
+
37
+ result = df.select(
38
+ [
39
+ ((pl.col(c) - pl.col(c).mean()) / pl.col(c).std()).alias(c)
40
+ for c in df.columns
41
+ ]
42
+ )
43
+
44
+ if return_series:
45
+ return result.to_series(0)
46
+ return result
47
+
48
+
49
+ def winsorize(data, cutoff=None, replace_with_cutoff=True):
50
+ """Winsorize a Polars DataFrame/Series with the largest/lowest value not considered outlier.
51
+
52
+ Args:
53
+ data (pl.DataFrame | pl.Series): Data to winsorize.
54
+ cutoff (dict): A dictionary with keys `{'std': [low, high]}` or
55
+ `{'quantile': [low, high]}`.
56
+ replace_with_cutoff (bool): If True, replace outliers with the cutoff
57
+ value; if False, replace them with the closest existing values
58
+ (default: True).
59
+
60
+ Returns:
61
+ pl.DataFrame | pl.Series: Winsorized data, the same type as the input.
62
+ """
63
+ return _transform_outliers(
64
+ data, cutoff, replace_with_cutoff=replace_with_cutoff, method="winsorize"
65
+ )
66
+
67
+
68
+ def trim(data, cutoff=None):
69
+ """Trim a Polars DataFrame/Series by replacing outlier values with NaNs.
70
+
71
+ Args:
72
+ data (pl.DataFrame | pl.Series): Data to trim.
73
+ cutoff (dict): A dictionary with keys `{'std': [low, high]}` or
74
+ `{'quantile': [low, high]}`.
75
+
76
+ Returns:
77
+ pl.DataFrame | pl.Series: Trimmed data (outliers replaced with NaN), the
78
+ same type as the input.
79
+ """
80
+ return _transform_outliers(data, cutoff, replace_with_cutoff=None, method="trim")
81
+
82
+
83
+ def _transform_outliers(data, cutoff, replace_with_cutoff, method):
84
+ """Winsorize or trim outliers in a Polars DataFrame/Series (shared by `trim` and `winsorize`).
85
+
86
+ Args:
87
+ data (pl.DataFrame | pl.Series): Data to transform.
88
+ cutoff (dict): A dictionary with keys `{'std': [low, high]}` or
89
+ `{'quantile': [low, high]}`.
90
+ replace_with_cutoff (bool | None): For winsorizing, replace outliers with
91
+ the cutoff value (True) or with the closest existing value (False).
92
+ Ignored when trimming.
93
+ method (str): 'winsorize' or 'trim'.
94
+
95
+ Returns:
96
+ pl.DataFrame | pl.Series: Transformed data, the same type as the input.
97
+ """
98
+ return_series = False
99
+ if isinstance(data, pl.DataFrame):
100
+ df = data.clone()
101
+ elif isinstance(data, pl.Series):
102
+ df = pl.DataFrame({data.name or "0": data})
103
+ return_series = True
104
+ else:
105
+ raise ValueError("Data must be a Polars DataFrame or Series")
106
+
107
+ # Transform each column if a DataFrame, if Series just transform data
108
+ transformed_cols = []
109
+ for col in df.columns:
110
+ # Get the series for calculations
111
+ series = df[col]
112
+
113
+ # Calculate cutoff values
114
+ if isinstance(cutoff, dict):
115
+ if "quantile" in cutoff:
116
+ quantiles = cutoff["quantile"]
117
+ # Use numpy quantile to match pandas interpolation behavior
118
+ series_array = series.to_numpy()
119
+ lower_q = float(np.quantile(series_array, quantiles[0]))
120
+ upper_q = (
121
+ float(np.quantile(series_array, quantiles[1]))
122
+ if len(quantiles) > 1
123
+ else lower_q
124
+ )
125
+ elif "std" in cutoff:
126
+ mean_val = series.mean()
127
+ std_val = series.std()
128
+ lower_q = mean_val - std_val * cutoff["std"][0]
129
+ upper_q = mean_val + std_val * cutoff["std"][1]
130
+ else:
131
+ raise ValueError(
132
+ "cutoff must be a dictionary with quantile or std keys."
133
+ )
134
+
135
+ # If replace_with_cutoff is false, replace with true existing values closest to cutoff
136
+ if method == "winsorize" and not replace_with_cutoff:
137
+ filtered_lower = series.filter(series > lower_q)
138
+ filtered_upper = series.filter(series < upper_q)
139
+ if len(filtered_lower) > 0:
140
+ lower_q = filtered_lower.min()
141
+ if len(filtered_upper) > 0:
142
+ upper_q = filtered_upper.max()
143
+
144
+ # Apply transformation using Polars expressions with column reference
145
+ if method == "trim":
146
+ # Replace outliers with null (NaN)
147
+ transformed_expr = (
148
+ pl.when(pl.col(col) < lower_q)
149
+ .then(None)
150
+ .when(pl.col(col) > upper_q)
151
+ .then(None)
152
+ .otherwise(pl.col(col))
153
+ )
154
+ elif method == "winsorize":
155
+ # Replace outliers with cutoff values
156
+ transformed_expr = (
157
+ pl.when(pl.col(col) < lower_q)
158
+ .then(lower_q)
159
+ .when(pl.col(col) > upper_q)
160
+ .then(upper_q)
161
+ .otherwise(pl.col(col))
162
+ )
163
+ else:
164
+ raise ValueError(f"Unknown method: {method}")
165
+
166
+ transformed_cols.append(transformed_expr.alias(col))
167
+ else:
168
+ raise ValueError("cutoff must be a dictionary with quantile or std keys.")
169
+
170
+ # Use with_columns to update all columns at once
171
+ result_df = df.with_columns(transformed_cols)
172
+
173
+ # Return Series if input was Series, otherwise DataFrame
174
+ if return_series:
175
+ return result_df.to_series(0)
176
+ return result_df
177
+
178
+
179
+ def find_spikes(
180
+ data,
181
+ global_spike_cutoff=3,
182
+ diff_spike_cutoff=3,
183
+ *,
184
+ TR: float | None = None,
185
+ sampling_freq: float | None = None,
186
+ ):
187
+ """Identify spikes (motion artifacts, intensity outliers) in 4D fMRI data.
188
+
189
+ Args:
190
+ data (BrainData | nib.Nifti1Image): 4D functional data.
191
+ global_spike_cutoff (float | None): Cutoff in standard deviations for
192
+ spikes in the per-TR global mean signal; None skips this detector.
193
+ Defaults to 3.
194
+ diff_spike_cutoff (float | None): Cutoff in standard deviations for
195
+ spikes in the per-TR mean absolute frame-to-frame difference; None
196
+ skips this detector. Defaults to 3.
197
+ TR (float | None): Repetition time in seconds; sets the returned
198
+ DesignMatrix's `sampling_freq` for downstream `.append()` /
199
+ `.convolve()`. Pass at most one of `TR` and `sampling_freq`.
200
+ sampling_freq (float | None): Sampling frequency in Hz (1 / TR). See `TR`.
201
+
202
+ Returns:
203
+ DesignMatrix: One indicator column per detected spike TR, named
204
+ `.nl_global_spike{n}` / `.nl_diff_spike{n}` in the reserved namespace
205
+ for generated columns (see `RESERVED_PREFIX`) and pre-marked as
206
+ confounds. Row position is the time axis. A volume flagged by both
207
+ detectors yields identical one-hot columns, so only the
208
+ `.nl_global_spike*` one is kept. Without `TR` / `sampling_freq` the
209
+ result has `sampling_freq=None` and can still be appended to a
210
+ DesignMatrix that has one.
211
+ """
212
+
213
+ from nltools.data import BrainData
214
+
215
+ if (global_spike_cutoff is None) & (diff_spike_cutoff is None):
216
+ raise ValueError("Did not input any cutoffs to identify spikes in this data.")
217
+
218
+ if isinstance(data, BrainData):
219
+ # Avoid deepcopy overhead - just copy the data array
220
+ data_array = data.data.copy()
221
+ global_mn = np.mean(data_array, axis=1)
222
+ frame_diff = np.mean(np.abs(np.diff(data_array, axis=0)), axis=1)
223
+ elif isinstance(data, nib.Nifti1Image):
224
+ # Avoid deepcopy overhead - just copy the data array
225
+ data_array = data.get_fdata().copy()
226
+ if len(data_array.shape) > 3:
227
+ data_array = np.squeeze(data_array)
228
+ elif len(data_array.shape) < 3:
229
+ raise ValueError("nibabel instance does not appear to be 4D data.")
230
+ global_mn = np.mean(data_array, axis=(0, 1, 2))
231
+ frame_diff = np.mean(np.abs(np.diff(data_array, axis=3)), axis=(0, 1, 2))
232
+ else:
233
+ raise ValueError(
234
+ "Currently this function can only accomodate BrainData and nibabel instances"
235
+ )
236
+
237
+ if global_spike_cutoff is not None:
238
+ # Vectorize outlier detection - avoid np.append in loops
239
+ global_mean = np.mean(global_mn)
240
+ global_std = np.std(global_mn)
241
+ upper_threshold = global_mean + global_std * global_spike_cutoff
242
+ lower_threshold = global_mean - global_std * global_spike_cutoff
243
+ global_outliers = np.where(
244
+ (global_mn > upper_threshold) | (global_mn < lower_threshold)
245
+ )[0]
246
+
247
+ if diff_spike_cutoff is not None:
248
+ # Vectorize outlier detection - avoid np.append in loops
249
+ diff_mean = np.mean(frame_diff)
250
+ diff_std = np.std(frame_diff)
251
+ upper_threshold = diff_mean + diff_std * diff_spike_cutoff
252
+ lower_threshold = diff_mean - diff_std * diff_spike_cutoff
253
+ frame_outliers = np.where(
254
+ (frame_diff > upper_threshold) | (frame_diff < lower_threshold)
255
+ )[0]
256
+ # Build spike regressors using Polars. Row position is the time axis;
257
+ # no separate "TR" index column (pandas-era artifact, dropped in v0.6.0).
258
+ outlier_data: dict[str, list[int]] = {}
259
+
260
+ if global_spike_cutoff is not None:
261
+ for i, loc in enumerate(global_outliers):
262
+ col_name = f"global_spike{i + 1}"
263
+ col_values = [0] * len(global_mn)
264
+ col_values[int(loc)] = 1
265
+ outlier_data[col_name] = col_values
266
+
267
+ if diff_spike_cutoff is not None:
268
+ for i, loc in enumerate(frame_outliers):
269
+ col_name = f"diff_spike{i + 1}"
270
+ col_values = [0] * len(global_mn)
271
+ col_values[int(loc)] = 1
272
+ outlier_data[col_name] = col_values
273
+
274
+ # Two detectors, one timeline: the same TR can be flagged by both, and a
275
+ # one-hot column per detection would then be literally duplicated —
276
+ # straight duplicate columns, which .append(axis=1) refuses. Deduplicate
277
+ # on the flagged position, keeping the global detection so the tie-break
278
+ # is deterministic rather than insertion-ordered. (Only the name is at
279
+ # stake: the colliding columns are bitwise identical.)
280
+ global_prefix = "global_spike"
281
+ seen: dict[int, str] = {}
282
+ for name in [c for c in outlier_data if c.startswith(global_prefix)] + [
283
+ c for c in outlier_data if not c.startswith(global_prefix)
284
+ ]:
285
+ loc = outlier_data[name].index(1)
286
+ if loc in seen:
287
+ del outlier_data[name]
288
+ else:
289
+ seen[loc] = name
290
+
291
+ if TR is not None and sampling_freq is not None:
292
+ raise ValueError(
293
+ "find_spikes: pass exactly one of `TR` or `sampling_freq`, not both."
294
+ )
295
+ if TR is not None:
296
+ sampling_freq = 1.0 / TR
297
+
298
+ from nltools.data.designmatrix.utils import _design_from_generated
299
+
300
+ # No spikes is a normal outcome, not an error. Polars cannot express
301
+ # "n rows, 0 columns", so hand the row count over explicitly — otherwise
302
+ # the result reports 0 rows and downstream `.append()` rejects it for not
303
+ # matching the rest of the design.
304
+ return _design_from_generated(
305
+ pl.DataFrame(outlier_data),
306
+ sampling_freq=sampling_freq,
307
+ n_rows=None if outlier_data else len(global_mn),
308
+ )
@@ -0,0 +1,83 @@
1
+ """Ordinary least squares on plain numpy arrays.
2
+
3
+ `regress` fits `Y ~ X` and returns coefficients, standard errors,
4
+ t-statistics, p-values, degrees of freedom, and residuals as arrays. Use it for
5
+ quick regressions on tabular or behavioral data; for voxel-wise models on
6
+ imaging data use `BrainData.fit(model='glm')`, which adds masking, run
7
+ handling, and the modeling helpers.
8
+ """
9
+
10
+ from __future__ import annotations
11
+
12
+ import numpy as np
13
+ from scipy.stats import t as t_dist
14
+
15
+
16
+ def regress(X, Y, *, stats: str = "full", tail: int | str = 2):
17
+ """Fit an OLS regression of `Y` on `X`.
18
+
19
+ Does not add an intercept; include one in `X` explicitly. If `Y` is 2D, a
20
+ separate regression is fit to each column.
21
+
22
+ Args:
23
+ X (np.ndarray): Design matrix, shape (n_samples, n_regressors).
24
+ Y (np.ndarray): Response, shape (n_samples,) or (n_samples, n_targets).
25
+ stats (str): 'full' returns the 6-tuple below, 'betas' returns just `b`,
26
+ 'tstats' returns `(b, t)`. Defaults to 'full'.
27
+ tail (int | str): 2 or 'two' for two-tailed p-values (default); 1 or
28
+ 'one' for a one-tailed test of beta > 0 (negate a regressor for the
29
+ other direction).
30
+
31
+ Returns:
32
+ tuple: `(b, se, t, p, df, res)` when `stats='full'`: coefficients,
33
+ standard errors, t-statistics, p-values (per `tail`), residual
34
+ degrees of freedom, and residuals. `stats='betas'` returns just `b`;
35
+ `stats='tstats'` returns `(b, t)`.
36
+ """
37
+ from .validation import _validate_tail_parameter
38
+
39
+ if stats not in ("full", "betas", "tstats"):
40
+ raise ValueError("stats must be one of 'full', 'betas', 'tstats'")
41
+ tail_internal = _validate_tail_parameter(tail)
42
+
43
+ X = np.asarray(X)
44
+ Y = np.asarray(Y)
45
+ y_was_1d = Y.ndim == 1
46
+ if y_was_1d:
47
+ Y = Y[:, np.newaxis]
48
+
49
+ b = np.linalg.pinv(X) @ Y # (n_regressors, n_targets)
50
+ if stats == "betas":
51
+ return b.squeeze()
52
+
53
+ res = Y - X @ b
54
+ df_scalar = X.shape[0] - X.shape[1]
55
+ # Unbiased residual SE from *uncentered* RSS: sqrt(RSS / (n - p)). Correct for
56
+ # both intercept and intercept-free models. np.std(res, ddof=p) would center
57
+ # the residuals, underestimating RSS when X has no intercept (a supported
58
+ # usage — see the docstring). Matches stats/correlation.py; see GH #287.
59
+ sigma = np.sqrt((res**2).sum(axis=0) / df_scalar) # (n_targets,)
60
+ xtx_inv_diag = np.diag(np.linalg.pinv(X.T @ X)) # (n_regressors,)
61
+ se = np.sqrt(xtx_inv_diag)[:, np.newaxis] * sigma[np.newaxis, :]
62
+
63
+ t = np.zeros_like(b)
64
+ mask = se > 1e-6
65
+ t[mask] = b[mask] / se[mask]
66
+
67
+ if stats == "tstats":
68
+ return b.squeeze(), t.squeeze()
69
+
70
+ df = np.full(t.shape[1], df_scalar)
71
+ if tail_internal == "upper":
72
+ p = 1 - t_dist.cdf(t, df)
73
+ else:
74
+ p = 2 * (1 - t_dist.cdf(np.abs(t), df))
75
+
76
+ return (
77
+ b.squeeze(),
78
+ se.squeeze(),
79
+ t.squeeze(),
80
+ p.squeeze(),
81
+ df.squeeze(),
82
+ res.squeeze(),
83
+ )
@@ -0,0 +1,303 @@
1
+ """Temporal signal processing — resampling, filtering, and basis functions."""
2
+
3
+ import numpy as np
4
+ import polars as pl
5
+ from scipy.interpolate import interp1d
6
+ from scipy.signal import butter, filtfilt
7
+
8
+
9
+ def calc_bpm(beat_interval, sampling_freq):
10
+ """Calculate instantaneous BPM from beat to beat interval.
11
+
12
+ Args:
13
+ beat_interval (int): Number of samples between beats (typically the R-R interval).
14
+ sampling_freq (float): Sampling frequency in Hz.
15
+
16
+ Returns:
17
+ float: Beats per minute for the time interval.
18
+ """
19
+ return 60 * sampling_freq * (1 / (beat_interval))
20
+
21
+
22
+ def downsample(
23
+ data, *, sampling_freq=None, target=None, target_type="samples", method="mean"
24
+ ):
25
+ """Downsample a Polars DataFrame/Series to a new target frequency or number of samples using averaging.
26
+
27
+ Args:
28
+ data (pl.DataFrame | pl.Series): Data to downsample.
29
+ sampling_freq (float): Sampling frequency of the data in Hz.
30
+ target (float): Downsampling target.
31
+ target_type (str): Unit of `target`, one of 'samples', 'seconds', or 'hz'.
32
+ Defaults to 'samples'.
33
+ method (str): Aggregation within each bin, 'mean' or 'median'. Defaults to 'mean'.
34
+
35
+ Returns:
36
+ pl.DataFrame | pl.Series: Downsampled data (same type as input).
37
+ """
38
+ if isinstance(data, pl.DataFrame):
39
+ df = data.clone()
40
+ return_series = False
41
+ elif isinstance(data, pl.Series):
42
+ df = pl.DataFrame({data.name or "0": data})
43
+ return_series = True
44
+ else:
45
+ raise ValueError("Data must be a Polars DataFrame or Series instance.")
46
+
47
+ if method not in ("mean", "median"):
48
+ raise ValueError("Metric must be either 'mean' or 'median'")
49
+
50
+ if target_type == "samples":
51
+ n_samples = target
52
+ elif target_type == "seconds":
53
+ n_samples = target * sampling_freq
54
+ elif target_type == "hz":
55
+ n_samples = sampling_freq / target
56
+ else:
57
+ raise ValueError('Make sure target_type is "samples", "seconds", or "hz".')
58
+
59
+ # Calculate grouping indices more efficiently (matches design_matrix.py pattern)
60
+ n_groups = int(np.ceil(df.shape[0] / n_samples))
61
+ idx = pl.Series(np.repeat(np.arange(n_groups), int(n_samples))[: df.shape[0]])
62
+
63
+ # Handle remainder samples (last incomplete group)
64
+ if df.shape[0] > len(idx):
65
+ remainder = pl.Series(np.repeat(idx[-1] + 1, df.shape[0] - len(idx)))
66
+ idx = pl.concat([idx, remainder])
67
+
68
+ # Add grouping index to dataframe
69
+ df_with_idx = df.with_columns(idx.alias("_group_idx"))
70
+
71
+ # Group by index and aggregate using Polars group_by
72
+ if method == "mean":
73
+ downsampled_df = (
74
+ df_with_idx.group_by("_group_idx", maintain_order=True)
75
+ .agg([pl.col(col).mean() for col in df.columns])
76
+ .drop("_group_idx")
77
+ )
78
+ else: # median
79
+ downsampled_df = (
80
+ df_with_idx.group_by("_group_idx", maintain_order=True)
81
+ .agg([pl.col(col).median() for col in df.columns])
82
+ .drop("_group_idx")
83
+ )
84
+
85
+ # Return Series if input was Series, otherwise DataFrame
86
+ if return_series:
87
+ return downsampled_df.to_series(0)
88
+ return downsampled_df
89
+
90
+
91
+ def upsample(
92
+ data, *, sampling_freq=None, target=None, target_type="samples", method="linear"
93
+ ):
94
+ """Upsample a Polars DataFrame/Series to a new target frequency or number of samples using interpolation.
95
+
96
+ Args:
97
+ data (pl.DataFrame | pl.Series): Data to upsample. Non-numeric columns
98
+ are dropped from a DataFrame.
99
+ sampling_freq (float): Sampling frequency of the data in Hz.
100
+ target (float): Upsampling target.
101
+ target_type (str): Unit of `target`, one of 'samples', 'seconds', or 'hz'.
102
+ method (str): Interpolation method, one of 'linear', 'nearest', 'zero',
103
+ 'slinear', 'quadratic', or 'cubic'; 'zero', 'slinear', 'quadratic'
104
+ and 'cubic' refer to spline interpolation of zeroth, first, second
105
+ or third order (default: 'linear').
106
+
107
+ Returns:
108
+ pl.DataFrame | pl.Series: Upsampled data, the same type as the input.
109
+ """
110
+ if isinstance(data, pl.DataFrame):
111
+ df = data.clone()
112
+ return_series = False
113
+ elif isinstance(data, pl.Series):
114
+ df = pl.DataFrame({data.name or "0": data})
115
+ return_series = True
116
+ else:
117
+ raise ValueError("Data must be a Polars DataFrame or Series instance.")
118
+
119
+ methods = ["linear", "nearest", "zero", "slinear", "quadratic", "cubic"]
120
+ if method not in methods:
121
+ raise ValueError(
122
+ "Method must be 'linear', 'nearest', 'zero', 'slinear', 'quadratic', 'cubic'"
123
+ )
124
+
125
+ if target_type == "samples":
126
+ n_samples = target
127
+ elif target_type == "seconds":
128
+ n_samples = target * sampling_freq
129
+ elif target_type == "hz":
130
+ n_samples = float(sampling_freq) / float(target)
131
+ else:
132
+ raise ValueError('Make sure target_type is "samples", "seconds", or "hz".')
133
+
134
+ orig_spacing = np.arange(0, df.shape[0], 1)
135
+ new_spacing = np.arange(0, df.shape[0] - 1, n_samples)
136
+
137
+ # Interpolate each column using scipy (matches stats.upsample logic)
138
+ upsampled_data = {}
139
+ for col in df.columns:
140
+ col_data = df[col].to_numpy()
141
+
142
+ # Create interpolation function
143
+ interpolate = interp1d(orig_spacing, col_data, kind=method)
144
+
145
+ # Interpolate to new indices
146
+ upsampled_data[col] = interpolate(new_spacing)
147
+
148
+ # Create new Polars DataFrame
149
+ upsampled_df = pl.DataFrame(upsampled_data)
150
+
151
+ # Return Series if input was Series, otherwise DataFrame
152
+ if return_series:
153
+ return upsampled_df.to_series(0)
154
+ return upsampled_df
155
+
156
+
157
+ def make_cosine_basis(nsamples, sampling_freq, filter_length, unit_scale=True, drop=0):
158
+ """Create basis functions for a discrete cosine transform.
159
+
160
+ Based on the implementation in ``spm_filter`` and ``spm_dctmtx`` because
161
+ scipy DCT can only apply transforms but not return the basis functions. Like
162
+ SPM, this does not add a constant (i.e. intercept), but does retain the first
163
+ basis (i.e. sigmoidal/linear drift).
164
+
165
+ Args:
166
+ nsamples (int): Number of observations (e.g. TRs).
167
+ sampling_freq (float): Sampling frequency in Hz (i.e. 1 / TR).
168
+ filter_length (int): Filter length in seconds.
169
+ unit_scale (bool): Scale the basis functions to the range [-1, 1]. Defaults to True.
170
+ drop (int): Number of leading (slowest) bases to drop after the constant is
171
+ removed; `drop=2` removes the first two. Defaults to 0, which keeps
172
+ the linear/sigmoidal drift basis that SPM discards.
173
+
174
+ Returns:
175
+ np.ndarray: Basis matrix of shape (nsamples, n_bases).
176
+ """
177
+
178
+ # Figure out number of basis functions to create
179
+ order = int(np.trunc(2 * (nsamples * sampling_freq) / filter_length + 1))
180
+
181
+ n = np.arange(nsamples)
182
+
183
+ # Initialize basis function matrix
184
+ C = np.zeros((len(n), order))
185
+
186
+ # Add constant
187
+ C[:, 0] = np.ones(len(n)) / np.sqrt(nsamples)
188
+
189
+ # Insert higher order cosine basis functions (vectorized)
190
+ if order > 1:
191
+ # Vectorize: create index matrix for broadcasting
192
+ i_indices = np.arange(1, order)[:, np.newaxis] # (order-1, 1)
193
+ n_indices = n[np.newaxis, :] # (1, nsamples)
194
+ # Compute all cosine basis functions at once
195
+ C[:, 1:] = (
196
+ np.sqrt(2.0 / nsamples)
197
+ * np.cos(np.pi * (2 * n_indices + 1) * i_indices / (2 * nsamples))
198
+ ).T
199
+
200
+ # Drop intercept ala SPM
201
+ C = C[:, 1:]
202
+
203
+ if C.size == 0:
204
+ raise ValueError(
205
+ "Basis function creation failed! nsamples is too small for requested filter_length."
206
+ )
207
+
208
+ if unit_scale:
209
+ C *= 1.0 / C[0, 0]
210
+
211
+ C = C[:, drop:]
212
+
213
+ return C
214
+
215
+
216
+ def _butter_bandpass_filter(data, low_cut, high_cut, fs, axis=0, order=5):
217
+ """Apply a bandpass butterworth filter with zero-phase filtering.
218
+
219
+ Args:
220
+ data (np.ndarray): Signal(s) to filter.
221
+ low_cut (float): Lower cutoff frequency (high-pass edge) in Hz.
222
+ high_cut (float): Upper cutoff frequency (low-pass edge) in Hz.
223
+ fs (float): Sampling frequency in Hz.
224
+ axis (int): Axis along which to filter. Defaults to 0.
225
+ order (int): Butterworth filter order. Defaults to 5.
226
+
227
+ Returns:
228
+ np.ndarray: Bandpass-filtered data with the same shape as `data`.
229
+ """
230
+ nyq = 0.5 * fs
231
+ b, a = butter(order, [low_cut / nyq, high_cut / nyq], btype="band")
232
+ return filtfilt(b, a, data, axis=axis)
233
+
234
+
235
+ def _phase_mean_angle(phase_angles):
236
+ """Compute the circular mean of phase angles.
237
+
238
+ Follows Fisher, N. I. (1995), *Statistical Analysis of Circular Data*.
239
+
240
+ Args:
241
+ phase_angles (np.ndarray): A 1D array of angles, or a 2D array whose rows
242
+ are sets of angles (e.g. time points x subjects).
243
+
244
+ Returns:
245
+ np.ndarray: The mean angle; one value per row for 2D input.
246
+ """
247
+
248
+ axis = 0 if len(phase_angles.shape) == 1 else 1
249
+ return np.arctan2(
250
+ np.mean(np.sin(phase_angles), axis=axis),
251
+ np.mean(np.cos(phase_angles), axis=axis),
252
+ )
253
+
254
+
255
+ def _phase_vector_length(phase_angles):
256
+ """Compute the mean resultant vector length of phase angles.
257
+
258
+ Follows Fisher, N. I. (1995), *Statistical Analysis of Circular Data*.
259
+
260
+ Args:
261
+ phase_angles (np.ndarray): A 1D array of angles, or a 2D array whose rows
262
+ are sets of angles (e.g. time points x subjects).
263
+
264
+ Returns:
265
+ np.ndarray: Vector length in [0, 1]; one value per row for 2D input.
266
+ """
267
+
268
+ axis = 0 if len(phase_angles.shape) == 1 else 1
269
+ return np.float32(
270
+ np.sqrt(
271
+ np.mean(np.cos(phase_angles), axis=axis) ** 2
272
+ + np.mean(np.sin(phase_angles), axis=axis) ** 2
273
+ )
274
+ )
275
+
276
+
277
+ def _phase_rayleigh_p(phase_angles):
278
+ """Compute Rayleigh-test p-values for non-uniformity of phase angles.
279
+
280
+ Follows Fisher, N. I. (1995), *Statistical Analysis of Circular Data*.
281
+
282
+ Args:
283
+ phase_angles (np.ndarray): A 1D array of angles, or a 2D array whose rows
284
+ are sets of angles (e.g. time points x subjects).
285
+
286
+ Returns:
287
+ np.ndarray: Rayleigh p-values; one value per row for 2D input.
288
+
289
+ Note:
290
+ The test treats the angles in each set as independent, which
291
+ autocorrelated timeseries violate.
292
+ """
293
+
294
+ n = len(phase_angles) if len(phase_angles.shape) == 1 else phase_angles.shape[1]
295
+
296
+ Z = n * _phase_vector_length(phase_angles) ** 2
297
+ if n <= 50:
298
+ return np.exp(-1 * Z) * (
299
+ 1
300
+ + (2 * Z - Z**2) / (4 * n)
301
+ - (24 * Z - 132 * Z**2 + 76 * Z**3 - 9 * Z**4) / (288 * n**2)
302
+ )
303
+ return np.exp(-1 * Z)