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.
- nltools/__init__.py +55 -0
- nltools/algorithms/__init__.py +90 -0
- nltools/algorithms/alignment/__init__.py +21 -0
- nltools/algorithms/alignment/procrustes.py +565 -0
- nltools/algorithms/alignment/srm.py +758 -0
- nltools/algorithms/backends.py +1059 -0
- nltools/algorithms/corrections.py +177 -0
- nltools/algorithms/decoding.py +327 -0
- nltools/algorithms/inference/__init__.py +50 -0
- nltools/algorithms/inference/bootstrap.py +1386 -0
- nltools/algorithms/inference/correlation.py +373 -0
- nltools/algorithms/inference/intersubject.py +422 -0
- nltools/algorithms/inference/isc.py +1554 -0
- nltools/algorithms/inference/matrix.py +602 -0
- nltools/algorithms/inference/one_sample.py +288 -0
- nltools/algorithms/inference/random.py +122 -0
- nltools/algorithms/inference/timeseries.py +347 -0
- nltools/algorithms/inference/two_sample.py +212 -0
- nltools/algorithms/inference/utils.py +58 -0
- nltools/algorithms/inference/validation.py +282 -0
- nltools/algorithms/neighborhoods.py +207 -0
- nltools/algorithms/outliers.py +308 -0
- nltools/algorithms/regression.py +83 -0
- nltools/algorithms/signal.py +303 -0
- nltools/algorithms/similarity.py +234 -0
- nltools/algorithms/validation.py +151 -0
- nltools/cross_validation.py +72 -0
- nltools/data/__init__.py +30 -0
- nltools/data/adjacency/__init__.py +875 -0
- nltools/data/adjacency/io.py +111 -0
- nltools/data/adjacency/modeling.py +569 -0
- nltools/data/adjacency/plotting.py +174 -0
- nltools/data/adjacency/state.py +349 -0
- nltools/data/adjacency/stats.py +596 -0
- nltools/data/adjacency/utils.py +79 -0
- nltools/data/atlases/__init__.py +23 -0
- nltools/data/atlases/labeling.py +158 -0
- nltools/data/atlases/loading.py +76 -0
- nltools/data/atlases/registry.py +96 -0
- nltools/data/atlases/reporting.py +456 -0
- nltools/data/braindata/__init__.py +2170 -0
- nltools/data/braindata/analysis.py +1381 -0
- nltools/data/braindata/bootstrap.py +398 -0
- nltools/data/braindata/io.py +896 -0
- nltools/data/braindata/modeling.py +594 -0
- nltools/data/braindata/plotting.py +501 -0
- nltools/data/braindata/prediction.py +1250 -0
- nltools/data/braindata/utils.py +348 -0
- nltools/data/braindata/validation.py +197 -0
- nltools/data/braindata/viewer.js +266 -0
- nltools/data/braindata/viewer.py +770 -0
- nltools/data/combine.py +27 -0
- nltools/data/designmatrix/__init__.py +1032 -0
- nltools/data/designmatrix/append.py +518 -0
- nltools/data/designmatrix/diagnostics.py +248 -0
- nltools/data/designmatrix/io.py +356 -0
- nltools/data/designmatrix/plotting.py +291 -0
- nltools/data/designmatrix/regressors.py +463 -0
- nltools/data/designmatrix/transforms.py +200 -0
- nltools/data/designmatrix/utils.py +350 -0
- nltools/data/ownership.py +129 -0
- nltools/data/results.py +291 -0
- nltools/data/roc/__init__.py +398 -0
- nltools/data/simulator/__init__.py +927 -0
- nltools/data/simulator/haxby.py +124 -0
- nltools/data/validation.py +83 -0
- nltools/datasets.py +218 -0
- nltools/io/__init__.py +10 -0
- nltools/io/events.py +67 -0
- nltools/io/h5.py +246 -0
- nltools/mask.py +403 -0
- nltools/models/__init__.py +11 -0
- nltools/models/glm.py +543 -0
- nltools/models/results.py +49 -0
- nltools/models/ridge.py +1303 -0
- nltools/models/validation.py +26 -0
- nltools/plotting/__init__.py +32 -0
- nltools/plotting/adjacency.py +421 -0
- nltools/plotting/brain.py +669 -0
- nltools/plotting/decomposition.py +111 -0
- nltools/plotting/prediction.py +110 -0
- nltools/resources/covariates_example.csv +161 -0
- nltools/resources/onsets_example.csv +40 -0
- nltools/templates/__init__.py +51 -0
- nltools/templates/config.py +144 -0
- nltools/templates/fetch.py +260 -0
- nltools/templates/matching.py +183 -0
- nltools/templates/paths.py +106 -0
- nltools/templates/registry.py +25 -0
- nltools/utils.py +230 -0
- nltools/version.py +13 -0
- nltools-0.6.0.dev0.dist-info/METADATA +95 -0
- nltools-0.6.0.dev0.dist-info/RECORD +95 -0
- nltools-0.6.0.dev0.dist-info/WHEEL +4 -0
- nltools-0.6.0.dev0.dist-info/licenses/LICENSE +21 -0
|
@@ -0,0 +1,350 @@
|
|
|
1
|
+
"""Shared helpers for DesignMatrix submodules.
|
|
2
|
+
|
|
3
|
+
These are internal utilities used by the facade and submodules — not part of the
|
|
4
|
+
public API.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
import functools
|
|
10
|
+
from copy import deepcopy
|
|
11
|
+
import re
|
|
12
|
+
from typing import TYPE_CHECKING
|
|
13
|
+
|
|
14
|
+
import polars as pl
|
|
15
|
+
|
|
16
|
+
from nltools.data.ownership import _copy_frame
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
if TYPE_CHECKING:
|
|
20
|
+
from nltools.data.designmatrix import DesignMatrix
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
RESERVED_PREFIX = ".nl_"
|
|
24
|
+
"""Prefix marking a column name generated by nltools rather than the user.
|
|
25
|
+
|
|
26
|
+
Every column nltools invents — polynomial drift (`.nl_poly_0`), DCT cosine
|
|
27
|
+
bases (`.nl_cosine_1`), spike indicators (`.nl_global_spike1`), and the
|
|
28
|
+
run-separated variants produced by a multi-run `DesignMatrix.append` — carries
|
|
29
|
+
this prefix. Machinery that needs to recognize its own columns keys on the
|
|
30
|
+
prefix, never on heuristics over user-controlled names (underscore counts,
|
|
31
|
+
substring matches), so users may name their own regressors anything without
|
|
32
|
+
colliding with nltools internals.
|
|
33
|
+
"""
|
|
34
|
+
|
|
35
|
+
_RUN_SEPARATED_RE = re.compile(re.escape(RESERVED_PREFIX) + r"r(\d+)_(.+)")
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
def _reserved_name(base: str) -> str:
|
|
39
|
+
"""Build a generated column name inside the reserved namespace.
|
|
40
|
+
|
|
41
|
+
Args:
|
|
42
|
+
base: Name without the reserved prefix, e.g. ``'poly_0'``.
|
|
43
|
+
|
|
44
|
+
Returns:
|
|
45
|
+
str: `base` prefixed with `RESERVED_PREFIX`, idempotently — a name
|
|
46
|
+
that already carries the prefix is returned unchanged.
|
|
47
|
+
"""
|
|
48
|
+
return base if _is_reserved_name(base) else f"{RESERVED_PREFIX}{base}"
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
def _is_reserved_name(name: str) -> bool:
|
|
52
|
+
"""Return True if ``name`` is in the nltools-generated column namespace."""
|
|
53
|
+
return name.startswith(RESERVED_PREFIX)
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
def _strip_reserved_prefix(name: str) -> str:
|
|
57
|
+
"""Return ``name`` without its reserved prefix (a no-op if it has none)."""
|
|
58
|
+
return name[len(RESERVED_PREFIX) :] if _is_reserved_name(name) else name
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
def _run_separated_name(run_idx: int, name: str) -> str:
|
|
62
|
+
"""Build the run-separated variant of a column name.
|
|
63
|
+
|
|
64
|
+
Run separation is an nltools-generated naming decision, so the result
|
|
65
|
+
always lands in the reserved namespace regardless of whether the source
|
|
66
|
+
column was user-named (``motion_x`` → ``.nl_r0_motion_x``) or already
|
|
67
|
+
generated (``.nl_poly_0`` → ``.nl_r0_poly_0``; prefixes never stack).
|
|
68
|
+
|
|
69
|
+
Args:
|
|
70
|
+
run_idx: Zero-based run index.
|
|
71
|
+
name: Column name to separate.
|
|
72
|
+
|
|
73
|
+
Returns:
|
|
74
|
+
str: ``.nl_r{run_idx}_{base}``.
|
|
75
|
+
"""
|
|
76
|
+
return f"{RESERVED_PREFIX}r{run_idx}_{_strip_reserved_prefix(name)}"
|
|
77
|
+
|
|
78
|
+
|
|
79
|
+
def _parse_run_separated(name: str) -> tuple[int, str] | None:
|
|
80
|
+
"""Split a run-separated column name into its run index and base name.
|
|
81
|
+
|
|
82
|
+
Args:
|
|
83
|
+
name: Column name to parse.
|
|
84
|
+
|
|
85
|
+
Returns:
|
|
86
|
+
tuple[int, str] | None: `(run_idx, base)` for a run-separated name (e.g.
|
|
87
|
+
`'.nl_r1_poly_0'` → `(1, 'poly_0')`), else None.
|
|
88
|
+
"""
|
|
89
|
+
match = _RUN_SEPARATED_RE.fullmatch(name)
|
|
90
|
+
return (int(match.group(1)), match.group(2)) if match else None
|
|
91
|
+
|
|
92
|
+
|
|
93
|
+
# Row selectors operate on a temporary column for zero-column designs.
|
|
94
|
+
_ROW_SELECTION = frozenset({"head", "tail", "slice", "filter", "limit"})
|
|
95
|
+
_MUTATORS = frozenset({"insert_column", "replace_column", "drop_in_place", "extend"})
|
|
96
|
+
|
|
97
|
+
|
|
98
|
+
def _design_from_generated(
|
|
99
|
+
frame: pl.DataFrame,
|
|
100
|
+
*,
|
|
101
|
+
sampling_freq: float | None = None,
|
|
102
|
+
n_rows: int | None = None,
|
|
103
|
+
) -> DesignMatrix:
|
|
104
|
+
"""Build a DesignMatrix whose every column is an nltools-generated confound.
|
|
105
|
+
|
|
106
|
+
Callers outside this package (`find_spikes`) name their columns plainly and
|
|
107
|
+
hand the frame here; the reserved prefix is applied in the one package that
|
|
108
|
+
owns the namespace. Every renamed column is marked a confound.
|
|
109
|
+
|
|
110
|
+
Args:
|
|
111
|
+
frame (pl.DataFrame): Generated columns under their plain names.
|
|
112
|
+
sampling_freq (float | None): Sampling frequency in Hz, or None.
|
|
113
|
+
n_rows (int | None): Row count to record when `frame` has no columns.
|
|
114
|
+
|
|
115
|
+
Returns:
|
|
116
|
+
DesignMatrix: The design, with every column in the reserved namespace
|
|
117
|
+
and listed in `confounds`.
|
|
118
|
+
"""
|
|
119
|
+
from nltools.data.designmatrix import DesignMatrix
|
|
120
|
+
|
|
121
|
+
renamed = frame.rename({name: _reserved_name(name) for name in frame.columns})
|
|
122
|
+
return DesignMatrix(
|
|
123
|
+
renamed,
|
|
124
|
+
sampling_freq=sampling_freq,
|
|
125
|
+
confounds=list(renamed.columns),
|
|
126
|
+
n_rows=n_rows,
|
|
127
|
+
)
|
|
128
|
+
|
|
129
|
+
|
|
130
|
+
def _effective_frame(dm: DesignMatrix) -> pl.DataFrame:
|
|
131
|
+
"""Represent recorded observations during operations on a column-less frame."""
|
|
132
|
+
if dm.data.width == 0 and dm._n_rows is not None:
|
|
133
|
+
return pl.DataFrame({"": pl.repeat(None, dm.shape[0], eager=True)})
|
|
134
|
+
return dm.data
|
|
135
|
+
|
|
136
|
+
|
|
137
|
+
def _replacement_names(frame: pl.DataFrame, exprs, named_exprs) -> list[str]:
|
|
138
|
+
"""Ask Polars which columns the supplied expressions produce."""
|
|
139
|
+
return frame.lazy().select(*exprs, **named_exprs).collect_schema().names()
|
|
140
|
+
|
|
141
|
+
|
|
142
|
+
# Drift terms generated by add_poly / add_dct_basis, matched on the base name
|
|
143
|
+
# left after the run prefix is stripped (e.g. '.nl_r1_poly_0' -> 'poly_0').
|
|
144
|
+
_DRIFT_BASE_RE = re.compile(r"(?:poly|cosine)_\d+")
|
|
145
|
+
_INTERCEPT_BASE_RE = re.compile(r"(?:poly|cosine)_0")
|
|
146
|
+
|
|
147
|
+
|
|
148
|
+
def _is_generated_intercept(name: str) -> bool:
|
|
149
|
+
"""Return True if ``name`` is an intercept column nltools generated.
|
|
150
|
+
|
|
151
|
+
Covers the zeroth-order drift terms from `add_poly` / `add_dct_basis`
|
|
152
|
+
(``.nl_poly_0`` / ``.nl_cosine_0``) and their run-separated variants. Both
|
|
153
|
+
are all-ones columns, so anything computing a correlation matrix has to
|
|
154
|
+
drop them. Keyed on the reserved namespace: a user column is never an
|
|
155
|
+
intercept by this definition, however it happens to be named.
|
|
156
|
+
"""
|
|
157
|
+
if not _is_reserved_name(name):
|
|
158
|
+
return False
|
|
159
|
+
parsed = _parse_run_separated(name)
|
|
160
|
+
base = parsed[1] if parsed is not None else _strip_reserved_prefix(name)
|
|
161
|
+
return _INTERCEPT_BASE_RE.fullmatch(base) is not None
|
|
162
|
+
|
|
163
|
+
|
|
164
|
+
def _has_run_separated_drift(dm: DesignMatrix) -> bool:
|
|
165
|
+
"""Return True if ``dm`` carries per-run polynomial or cosine drift terms.
|
|
166
|
+
|
|
167
|
+
Adding a global drift term to a design that already models drift per run
|
|
168
|
+
is ambiguous, so both `add_poly` and `add_dct_basis` refuse it. Detection
|
|
169
|
+
keys on the reserved namespace nltools controls (``.nl_r{run}_poly_{i}`` /
|
|
170
|
+
``.nl_r{run}_cosine_{i}``), so a user confound is never mistaken for one
|
|
171
|
+
however it is named.
|
|
172
|
+
"""
|
|
173
|
+
for col in dm.confounds or []:
|
|
174
|
+
parsed = _parse_run_separated(col)
|
|
175
|
+
if parsed is not None and _DRIFT_BASE_RE.fullmatch(parsed[1]):
|
|
176
|
+
return True
|
|
177
|
+
return False
|
|
178
|
+
|
|
179
|
+
|
|
180
|
+
def _is_column_selection(value) -> bool:
|
|
181
|
+
"""Recognize plain Polars column selectors without interpreting expressions."""
|
|
182
|
+
if isinstance(value, str):
|
|
183
|
+
return True
|
|
184
|
+
if isinstance(value, pl.Expr):
|
|
185
|
+
return value.meta.is_column_selection()
|
|
186
|
+
if isinstance(value, (list, tuple)):
|
|
187
|
+
return all(_is_column_selection(item) for item in value)
|
|
188
|
+
return False
|
|
189
|
+
|
|
190
|
+
|
|
191
|
+
def _df_passthrough(dm: DesignMatrix, name: str):
|
|
192
|
+
"""Forward Polars operations with explicit row, column and mutation context."""
|
|
193
|
+
attr = getattr(dm.data, name)
|
|
194
|
+
if not callable(attr):
|
|
195
|
+
return attr
|
|
196
|
+
|
|
197
|
+
@functools.wraps(attr) # nosemgrep: kwargs-internal-forwarding # Polars adapter
|
|
198
|
+
def wrapped(*args, **kwargs):
|
|
199
|
+
mutation = name in _MUTATORS or (
|
|
200
|
+
name in {"hstack", "vstack", "shrink_to_fit"}
|
|
201
|
+
and kwargs.get("in_place", False)
|
|
202
|
+
)
|
|
203
|
+
frame = _copy_frame(dm.data) if mutation else dm.data
|
|
204
|
+
if name in _ROW_SELECTION:
|
|
205
|
+
frame = _effective_frame(dm)
|
|
206
|
+
result = getattr(frame, name)(*args, **kwargs)
|
|
207
|
+
if (
|
|
208
|
+
name in {"insert_column", "hstack"}
|
|
209
|
+
and dm.data.width == 0
|
|
210
|
+
and dm._n_rows is not None
|
|
211
|
+
):
|
|
212
|
+
populated = frame if mutation else result
|
|
213
|
+
if populated.height != dm.shape[0]:
|
|
214
|
+
raise ValueError(
|
|
215
|
+
"Added columns must match the recorded number of rows."
|
|
216
|
+
)
|
|
217
|
+
operation = "unknown"
|
|
218
|
+
rename = None
|
|
219
|
+
replaced = None
|
|
220
|
+
if name in _ROW_SELECTION:
|
|
221
|
+
operation = "preserve"
|
|
222
|
+
elif (
|
|
223
|
+
name == "select"
|
|
224
|
+
and not kwargs
|
|
225
|
+
and all(_is_column_selection(arg) for arg in args)
|
|
226
|
+
):
|
|
227
|
+
operation = "preserve"
|
|
228
|
+
elif name == "rename":
|
|
229
|
+
operation = "rename"
|
|
230
|
+
rename = dict(zip(dm.columns, result.columns))
|
|
231
|
+
elif name == "replace_column":
|
|
232
|
+
operation = "replace"
|
|
233
|
+
index = args[0] if args else kwargs["index"]
|
|
234
|
+
replaced = [dm.columns[index], frame.columns[index]]
|
|
235
|
+
elif name in {"insert_column", "hstack"}:
|
|
236
|
+
operation = "preserve"
|
|
237
|
+
elif name in {"drop_in_place", "shrink_to_fit"}:
|
|
238
|
+
operation = "preserve"
|
|
239
|
+
if mutation:
|
|
240
|
+
updated = _copy_with(dm, frame, operation=operation, replaced=replaced)
|
|
241
|
+
dm.__dict__.update(updated.__dict__)
|
|
242
|
+
return dm if result is frame else result
|
|
243
|
+
if isinstance(result, pl.DataFrame):
|
|
244
|
+
n_rows = result.height
|
|
245
|
+
if name in _ROW_SELECTION and dm.data.width == 0:
|
|
246
|
+
result = pl.DataFrame()
|
|
247
|
+
elif name == "select" and operation == "preserve" and result.width == 0:
|
|
248
|
+
n_rows = dm.shape[0]
|
|
249
|
+
return _copy_with(
|
|
250
|
+
dm, result, operation=operation, rename=rename, n_rows=n_rows
|
|
251
|
+
)
|
|
252
|
+
if isinstance(result, pl.Series):
|
|
253
|
+
return _copy_frame(result.to_frame()).to_series()
|
|
254
|
+
return result
|
|
255
|
+
|
|
256
|
+
return wrapped
|
|
257
|
+
|
|
258
|
+
|
|
259
|
+
def _copy_with(
|
|
260
|
+
dm: DesignMatrix,
|
|
261
|
+
new_df: pl.DataFrame,
|
|
262
|
+
*,
|
|
263
|
+
operation: str = "preserve",
|
|
264
|
+
rename: dict | None = None,
|
|
265
|
+
replaced: list[str] | None = None,
|
|
266
|
+
sampling_freq=...,
|
|
267
|
+
convolved: list[str] | None = None,
|
|
268
|
+
confounds: list[str] | None = None,
|
|
269
|
+
multi: bool | None = None,
|
|
270
|
+
n_rows: int | None = None,
|
|
271
|
+
run_count: int | None = None,
|
|
272
|
+
) -> DesignMatrix:
|
|
273
|
+
"""Own a transformed frame and apply its caller-established metadata policy."""
|
|
274
|
+
from nltools.data.designmatrix import DesignMatrix
|
|
275
|
+
|
|
276
|
+
metadata = _get_metadata(dm)
|
|
277
|
+
if operation == "unknown":
|
|
278
|
+
metadata.update(
|
|
279
|
+
sampling_freq=None, convolved=[], confounds=[], multi=False, run_count=0
|
|
280
|
+
)
|
|
281
|
+
elif operation == "rename":
|
|
282
|
+
for key in ("convolved", "confounds"):
|
|
283
|
+
metadata[key] = [(rename or {}).get(c, c) for c in metadata[key]]
|
|
284
|
+
elif operation == "replace":
|
|
285
|
+
metadata["convolved"] = [
|
|
286
|
+
c for c in metadata["convolved"] if c not in (replaced or [])
|
|
287
|
+
]
|
|
288
|
+
if sampling_freq is not ...:
|
|
289
|
+
metadata["sampling_freq"] = sampling_freq
|
|
290
|
+
if convolved is not None:
|
|
291
|
+
metadata["convolved"] = convolved
|
|
292
|
+
if confounds is not None:
|
|
293
|
+
metadata["confounds"] = confounds
|
|
294
|
+
if multi is not None:
|
|
295
|
+
metadata["multi"] = multi
|
|
296
|
+
if run_count is not None:
|
|
297
|
+
metadata["run_count"] = run_count
|
|
298
|
+
for key in ("convolved", "confounds"):
|
|
299
|
+
metadata[key] = [c for c in metadata[key] if c in new_df.columns]
|
|
300
|
+
new = DesignMatrix.__new__(DesignMatrix)
|
|
301
|
+
memo = {id(dm): new}
|
|
302
|
+
new.data = _copy_frame(new_df, memo)
|
|
303
|
+
new.sampling_freq = metadata["sampling_freq"]
|
|
304
|
+
new._convolved = deepcopy(metadata["convolved"], memo)
|
|
305
|
+
new._confounds = deepcopy(metadata["confounds"], memo)
|
|
306
|
+
new.multi = metadata["multi"]
|
|
307
|
+
new._run_count = metadata["run_count"]
|
|
308
|
+
new._n_rows = (
|
|
309
|
+
(dm.shape[0] if n_rows is None else n_rows) if new_df.width == 0 else None
|
|
310
|
+
)
|
|
311
|
+
return new
|
|
312
|
+
|
|
313
|
+
|
|
314
|
+
def _get_metadata(dm: DesignMatrix) -> dict:
|
|
315
|
+
"""Extract metadata as a dict (for copying).
|
|
316
|
+
|
|
317
|
+
Args:
|
|
318
|
+
dm (DesignMatrix): DesignMatrix instance.
|
|
319
|
+
|
|
320
|
+
Returns:
|
|
321
|
+
dict: Dictionary with keys 'sampling_freq', 'convolved', 'confounds',
|
|
322
|
+
'multi', 'n_rows'.
|
|
323
|
+
"""
|
|
324
|
+
return {
|
|
325
|
+
"sampling_freq": dm.sampling_freq,
|
|
326
|
+
"convolved": dm.convolved.copy(),
|
|
327
|
+
"confounds": dm.confounds.copy(),
|
|
328
|
+
"multi": dm.multi,
|
|
329
|
+
"n_rows": dm._n_rows,
|
|
330
|
+
"run_count": dm._run_count,
|
|
331
|
+
}
|
|
332
|
+
|
|
333
|
+
|
|
334
|
+
def _get_data_columns(dm: DesignMatrix, exclude_confounds: bool = True) -> list[str]:
|
|
335
|
+
"""Get column names, optionally excluding confound regressors.
|
|
336
|
+
|
|
337
|
+
Used wherever experimental regressors must be distinguished from
|
|
338
|
+
nuisance/confound columns (polynomial drift, DCT cosines, motion, etc.).
|
|
339
|
+
|
|
340
|
+
Args:
|
|
341
|
+
dm (DesignMatrix): DesignMatrix instance.
|
|
342
|
+
exclude_confounds (bool): If True, exclude nuisance columns tracked in
|
|
343
|
+
``dm.confounds`` from the result. Default: True.
|
|
344
|
+
|
|
345
|
+
Returns:
|
|
346
|
+
list[str]: Column names (excluding confounds if requested).
|
|
347
|
+
"""
|
|
348
|
+
if exclude_confounds and dm.confounds:
|
|
349
|
+
return [col for col in dm.columns if col not in dm.confounds]
|
|
350
|
+
return list(dm.columns)
|
|
@@ -0,0 +1,129 @@
|
|
|
1
|
+
"""Copying that leaves the clone owning its own buffers.
|
|
2
|
+
|
|
3
|
+
A data class holds polars frames whose storage may be a view onto a NumPy
|
|
4
|
+
array the user still holds, and frames whose `pl.Object` cells are Python
|
|
5
|
+
objects shared with the original. Copying either naively hands the clone a
|
|
6
|
+
buffer or a cell somebody else can mutate. `_copy_frame` detaches one frame;
|
|
7
|
+
`_copy_graph` walks a whole object's `__dict__` and detaches every frame it
|
|
8
|
+
finds, preserving the aliases inside that graph through a shared memo.
|
|
9
|
+
|
|
10
|
+
The two frame copiers are deliberately not the same function: `_copy_frame`
|
|
11
|
+
re-`gather`s every non-Object series so a `DesignMatrix` clone owns its
|
|
12
|
+
numeric buffers outright, while `_copy_object_frames` clones the frame and
|
|
13
|
+
rewrites only its `pl.Object` columns, which is what `BrainData` and
|
|
14
|
+
`Adjacency` metadata need.
|
|
15
|
+
"""
|
|
16
|
+
|
|
17
|
+
from copy import deepcopy
|
|
18
|
+
|
|
19
|
+
import polars as pl
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def _copy_frame(frame: pl.DataFrame, memo: dict | None = None) -> pl.DataFrame:
|
|
23
|
+
"""Detach frame storage and Python Object cells with a shared copy memo."""
|
|
24
|
+
if memo is None:
|
|
25
|
+
memo = {}
|
|
26
|
+
if id(frame) in memo:
|
|
27
|
+
return memo[id(frame)]
|
|
28
|
+
frames = []
|
|
29
|
+
visited = set()
|
|
30
|
+
|
|
31
|
+
def discover(value):
|
|
32
|
+
if id(value) in visited or id(value) in memo:
|
|
33
|
+
return
|
|
34
|
+
visited.add(id(value))
|
|
35
|
+
if isinstance(value, pl.DataFrame):
|
|
36
|
+
frames.append(value)
|
|
37
|
+
for series in value:
|
|
38
|
+
if series.dtype == pl.Object:
|
|
39
|
+
for cell in series:
|
|
40
|
+
discover(cell)
|
|
41
|
+
elif isinstance(value, dict):
|
|
42
|
+
for key, item in value.items():
|
|
43
|
+
discover(key)
|
|
44
|
+
discover(item)
|
|
45
|
+
elif isinstance(value, (list, tuple)):
|
|
46
|
+
for item in value:
|
|
47
|
+
discover(item)
|
|
48
|
+
|
|
49
|
+
discover(frame)
|
|
50
|
+
for source in frames:
|
|
51
|
+
memo[id(source)] = source.clone()
|
|
52
|
+
for source in frames:
|
|
53
|
+
for index, series in enumerate(source):
|
|
54
|
+
if series.dtype == pl.Object:
|
|
55
|
+
detached = pl.Series(
|
|
56
|
+
series.name,
|
|
57
|
+
[deepcopy(cell, memo) for cell in series],
|
|
58
|
+
dtype=pl.Object,
|
|
59
|
+
)
|
|
60
|
+
else:
|
|
61
|
+
# Gather owns buffers even when the frame wraps a NumPy view.
|
|
62
|
+
detached = series.gather(pl.int_range(0, len(series), eager=True))
|
|
63
|
+
memo[id(source)].replace_column(index, detached)
|
|
64
|
+
return memo[id(frame)]
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
def _copy_graph(source, *, memo=None, exclude=(), replacements=None):
|
|
68
|
+
"""Copy one retained object graph, preserving its internal aliases."""
|
|
69
|
+
if memo is None:
|
|
70
|
+
memo = {}
|
|
71
|
+
if id(source) in memo:
|
|
72
|
+
return memo[id(source)]
|
|
73
|
+
new = type(source).__new__(type(source))
|
|
74
|
+
memo[id(source)] = new
|
|
75
|
+
values = {
|
|
76
|
+
key: value for key, value in source.__dict__.items() if key not in exclude
|
|
77
|
+
}
|
|
78
|
+
if replacements is not None:
|
|
79
|
+
values.update(replacements)
|
|
80
|
+
_copy_object_frames(values, memo)
|
|
81
|
+
for key, value in values.items():
|
|
82
|
+
setattr(new, key, deepcopy(value, memo))
|
|
83
|
+
return new
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
def _copy_object_frames(values, memo):
|
|
87
|
+
"""Prepare Polars Object cells for deepcopy without sharing Python objects."""
|
|
88
|
+
frames = []
|
|
89
|
+
seen = set()
|
|
90
|
+
|
|
91
|
+
def discover(value):
|
|
92
|
+
if id(value) in seen or id(value) in memo:
|
|
93
|
+
return
|
|
94
|
+
seen.add(id(value))
|
|
95
|
+
if isinstance(value, pl.DataFrame):
|
|
96
|
+
frames.append(value)
|
|
97
|
+
for series in value:
|
|
98
|
+
if series.dtype == pl.Object:
|
|
99
|
+
for cell in series:
|
|
100
|
+
discover(cell)
|
|
101
|
+
elif isinstance(value, dict):
|
|
102
|
+
for key, item in value.items():
|
|
103
|
+
discover(key)
|
|
104
|
+
discover(item)
|
|
105
|
+
elif isinstance(value, (list, tuple)):
|
|
106
|
+
for item in value:
|
|
107
|
+
discover(item)
|
|
108
|
+
|
|
109
|
+
discover(values)
|
|
110
|
+
# Register all frames first, including frames referred to by Object cells.
|
|
111
|
+
# The common memo preserves cycles and cell aliases across metadata frames.
|
|
112
|
+
for frame in frames:
|
|
113
|
+
memo[id(frame)] = frame.clone()
|
|
114
|
+
for frame in frames:
|
|
115
|
+
for index, series in enumerate(frame):
|
|
116
|
+
if series.dtype == pl.Object:
|
|
117
|
+
memo[id(frame)].replace_column(
|
|
118
|
+
index,
|
|
119
|
+
pl.Series(
|
|
120
|
+
series.name,
|
|
121
|
+
[deepcopy(cell, memo) for cell in series],
|
|
122
|
+
dtype=pl.Object,
|
|
123
|
+
),
|
|
124
|
+
)
|
|
125
|
+
|
|
126
|
+
|
|
127
|
+
def _copy_complete(source, memo=None):
|
|
128
|
+
"""Return a complete independently owned snapshot."""
|
|
129
|
+
return _copy_graph(source, memo=memo)
|