nltools 0.6.0.dev1__tar.gz → 0.6.0.dev3__tar.gz
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-0.6.0.dev1 → nltools-0.6.0.dev3}/PKG-INFO +1 -1
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/algorithms/alignment/procrustes.py +72 -76
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/algorithms/alignment/srm.py +78 -48
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/algorithms/backends.py +31 -9
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/algorithms/corrections.py +71 -18
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/algorithms/decoding.py +4 -2
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/algorithms/inference/bootstrap.py +55 -7
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/algorithms/inference/correlation.py +13 -5
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/algorithms/inference/intersubject.py +27 -27
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/algorithms/inference/isc.py +359 -229
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/algorithms/inference/one_sample.py +6 -2
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/algorithms/inference/timeseries.py +6 -3
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/algorithms/inference/two_sample.py +6 -2
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/algorithms/neighborhoods.py +7 -1
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/algorithms/outliers.py +52 -10
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/algorithms/regression.py +27 -8
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/algorithms/signal.py +103 -31
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/algorithms/similarity.py +10 -10
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/cross_validation.py +12 -2
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/adjacency/__init__.py +5 -1
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/adjacency/io.py +1 -1
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/adjacency/modeling.py +31 -13
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/adjacency/plotting.py +10 -5
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/adjacency/stats.py +14 -5
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/atlases/reporting.py +18 -15
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/braindata/__init__.py +98 -39
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/braindata/analysis.py +155 -89
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/braindata/bootstrap.py +24 -1
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/braindata/io.py +288 -98
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/braindata/modeling.py +4 -1
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/braindata/plotting.py +32 -3
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/braindata/prediction.py +19 -5
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/braindata/utils.py +132 -13
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/braindata/validation.py +39 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/designmatrix/__init__.py +10 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/designmatrix/diagnostics.py +6 -2
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/designmatrix/regressors.py +46 -6
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/designmatrix/transforms.py +39 -34
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/ownership.py +22 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/results.py +5 -3
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/results_io.py +9 -1
- nltools-0.6.0.dev3/nltools/data/roc/__init__.py +604 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/simulator/__init__.py +74 -50
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/simulator/haxby.py +29 -4
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/validation.py +11 -10
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/datasets.py +15 -5
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/io/h5.py +26 -20
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/mask.py +26 -32
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/models/ridge.py +56 -26
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/plotting/adjacency.py +9 -4
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/plotting/brain.py +22 -23
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/plotting/decomposition.py +4 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/plotting/prediction.py +2 -2
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/templates/matching.py +8 -5
- nltools-0.6.0.dev3/nltools/tests/core/test_algorithms/test_corrections.py +127 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/test_algorithms/test_decoding.py +13 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/test_algorithms/test_intersubject.py +29 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/test_algorithms/test_neighborhoods.py +15 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/test_algorithms/test_outliers.py +60 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/test_algorithms/test_procrustes.py +79 -301
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/test_algorithms/test_regression.py +47 -0
- nltools-0.6.0.dev3/nltools/tests/core/test_algorithms/test_signal.py +127 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/test_algorithms/test_similarity.py +48 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/test_backends.py +28 -23
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/test_bootstrap.py +73 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/test_cross_validation.py +24 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/test_inference/test_correlation.py +29 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/test_inference/test_cpu_parallelization.py +5 -41
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/test_inference/test_isc_group.py +53 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/test_inference/test_isc_vocabulary.py +22 -2
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/test_inference/test_one_sample.py +12 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/test_inference/test_timeseries.py +44 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/test_inference/test_two_sample.py +12 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/test_isc.py +96 -37
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/test_mask.py +81 -1
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/test_srm.py +25 -1
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/datasets/test_datasets.py +21 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/io_tests/test_h5.py +54 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/models/test_ridge.py +11 -5
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/plotting/test_adjacency.py +110 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/plotting/test_surface.py +49 -0
- nltools-0.6.0.dev3/nltools/tests/support/test_scripts.py +105 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/templates/test_brainspace.py +12 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/pyproject.toml +1 -1
- nltools-0.6.0.dev1/nltools/data/roc/__init__.py +0 -398
- nltools-0.6.0.dev1/nltools/tests/core/test_algorithms/test_corrections.py +0 -67
- nltools-0.6.0.dev1/nltools/tests/core/test_algorithms/test_signal.py +0 -72
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/.gitignore +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/LICENSE +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/README.md +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/__init__.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/algorithms/__init__.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/algorithms/alignment/__init__.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/algorithms/inference/__init__.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/algorithms/inference/matrix.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/algorithms/inference/random.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/algorithms/inference/utils.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/algorithms/inference/validation.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/algorithms/validation.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/__init__.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/adjacency/state.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/adjacency/utils.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/atlases/__init__.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/atlases/labeling.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/atlases/loading.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/atlases/registry.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/braindata/viewer.js +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/braindata/viewer.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/combine.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/designmatrix/append.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/designmatrix/io.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/designmatrix/plotting.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/data/designmatrix/utils.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/io/__init__.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/io/events.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/models/__init__.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/models/glm.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/models/results.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/models/validation.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/plotting/__init__.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/resources/covariates_example.csv +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/resources/onsets_example.csv +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/templates/__init__.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/templates/config.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/templates/fetch.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/templates/paths.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/templates/registry.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/conftest.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/__init__.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/test_algorithms/__init__.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/test_algorithms/conftest.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/test_gpu_policy.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/test_hyperalignment.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/test_inference/__init__.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/test_inference/test_api_conventions.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/test_inference/test_matrix.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/test_inference/test_progress_bar.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/test_inference/test_tail_vocabulary.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/test_inference/test_utils.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/core/test_utils.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/datasets/__init__.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/io_tests/__init__.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/io_tests/test_file_reader.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/models/__init__.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/models/conftest.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/models/test_glm.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/models/test_glm_warnings.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/models/test_results.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/plotting/__init__.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/plotting/test_f123_prediction.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/pyodide/.gitignore +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/pyodide/test_runner.mjs +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/support/__init__.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/support/test_designation.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/templates/__init__.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/templates/test_fetch_pyodide.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/utils/__init__.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/tests/utils/test_utils.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/utils.py +0 -0
- {nltools-0.6.0.dev1 → nltools-0.6.0.dev3}/nltools/version.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.5
|
|
2
2
|
Name: nltools
|
|
3
|
-
Version: 0.6.0.
|
|
3
|
+
Version: 0.6.0.dev3
|
|
4
4
|
Summary: A Python package to analyze neuroimaging data
|
|
5
5
|
Project-URL: Homepage, https://nltools.org
|
|
6
6
|
Author-email: "Luke J. Chang" <luke.j.chang@dartmouth.edu>, Eshin Jolly <eshin.jolly@gmail.com>
|
|
@@ -109,6 +109,23 @@ def _hyperalign(data, n_iter):
|
|
|
109
109
|
return aligned, transformation_matrix, template, disparity, scale
|
|
110
110
|
|
|
111
111
|
|
|
112
|
+
def _procrustes_similarity(mat1, mat2):
|
|
113
|
+
"""One minus scipy's Procrustes disparity, the only value inference needs.
|
|
114
|
+
|
|
115
|
+
Module level so `loky` can pickle it, and returning the scalar rather than
|
|
116
|
+
scipy's full tuple keeps the two transformed matrices out of the payload
|
|
117
|
+
every permutation ships back to the parent process.
|
|
118
|
+
|
|
119
|
+
Args:
|
|
120
|
+
mat1 (np.ndarray): Reference matrix, `(n_rows, n_cols)`.
|
|
121
|
+
mat2 (np.ndarray): Matrix fit to `mat1`, same shape.
|
|
122
|
+
|
|
123
|
+
Returns:
|
|
124
|
+
float: `1 - disparity`, higher meaning more similar.
|
|
125
|
+
"""
|
|
126
|
+
return 1 - procrust(mat1, mat2)[2]
|
|
127
|
+
|
|
128
|
+
|
|
112
129
|
def align(
|
|
113
130
|
data,
|
|
114
131
|
method="deterministic_srm",
|
|
@@ -134,7 +151,10 @@ def align(
|
|
|
134
151
|
method (str): One of `'probabilistic_srm'`, `'deterministic_srm'`, or
|
|
135
152
|
`'procrustes'`. Defaults to `'deterministic_srm'`.
|
|
136
153
|
n_features (int | None): Number of features in the common space (SRM only).
|
|
137
|
-
None uses the
|
|
154
|
+
None uses the smallest subject's size along the aligned axis (voxels
|
|
155
|
+
for `axis=0`), the largest value every subject can support; an
|
|
156
|
+
explicit value may not exceed that size for any subject. Must be None
|
|
157
|
+
for `'procrustes'`.
|
|
138
158
|
axis (int): Axis to align on: 0 aligns timepoints (ISC computed per voxel),
|
|
139
159
|
1 aligns voxels (ISC computed per timepoint). Defaults to 0.
|
|
140
160
|
n_iter (int): Number of `_SRM`/`_DetSRM` iterations; ignored by
|
|
@@ -153,9 +173,14 @@ def align(
|
|
|
153
173
|
|
|
154
174
|
Raises:
|
|
155
175
|
ValueError: If `data` is not a same-typed list, `method` or `axis` is
|
|
156
|
-
unknown,
|
|
176
|
+
unknown, `method='procrustes'` is combined with `axis=1` on
|
|
157
177
|
`BrainData` input — that transform spans images on both axes and has
|
|
158
|
-
no voxel axis to be returned on
|
|
178
|
+
no voxel axis to be returned on — or `method='procrustes'` is given
|
|
179
|
+
`BrainData` subjects with different voxel counts, whose zero-padded
|
|
180
|
+
results would not fit their own masks. Pass the subjects' `.data`
|
|
181
|
+
arrays to get the zero-padded result instead; it has no mask that
|
|
182
|
+
could describe it. An SRM `n_features` above any subject's voxel
|
|
183
|
+
count also raises.
|
|
159
184
|
|
|
160
185
|
Examples:
|
|
161
186
|
```python
|
|
@@ -173,7 +198,7 @@ def align(
|
|
|
173
198
|
```
|
|
174
199
|
"""
|
|
175
200
|
|
|
176
|
-
from nltools.data import BrainData
|
|
201
|
+
from nltools.data import BrainData
|
|
177
202
|
|
|
178
203
|
if not isinstance(data, list):
|
|
179
204
|
raise ValueError("Make sure you are inputting data is a list.")
|
|
@@ -185,6 +210,7 @@ def align(
|
|
|
185
210
|
)
|
|
186
211
|
|
|
187
212
|
if isinstance(data[0], BrainData):
|
|
213
|
+
from nltools.data.braindata.analysis import _brain_result
|
|
188
214
|
from nltools.data.braindata.utils import _result_from_array
|
|
189
215
|
|
|
190
216
|
data_type = "BrainData"
|
|
@@ -211,7 +237,7 @@ def align(
|
|
|
211
237
|
out = {}
|
|
212
238
|
if method in ["deterministic_srm", "probabilistic_srm"]:
|
|
213
239
|
if n_features is None:
|
|
214
|
-
n_features = int(
|
|
240
|
+
n_features = int(min(x.shape[0] for x in data))
|
|
215
241
|
if method == "deterministic_srm":
|
|
216
242
|
srm = _DetSRM(
|
|
217
243
|
n_features=n_features, n_iter=n_iter, random_state=random_state
|
|
@@ -251,18 +277,22 @@ def align(
|
|
|
251
277
|
|
|
252
278
|
if data_type == "BrainData":
|
|
253
279
|
if method == "procrustes":
|
|
280
|
+
# `_hyperalign` zero-pads every subject's feature axis up to the
|
|
281
|
+
# widest subject, so a narrower subject's result is wider than its
|
|
282
|
+
# own mask. `_brain_result` refuses that rather than returning an
|
|
283
|
+
# object whose `to_nifti` fails later.
|
|
254
284
|
out["transformed"] = [
|
|
255
|
-
|
|
285
|
+
_brain_result(source, values.T, "transformed", rows="preserve")
|
|
256
286
|
for source, values in zip(sources, out["transformed"])
|
|
257
287
|
]
|
|
258
|
-
out["common_model"] =
|
|
259
|
-
sources[0], out["common_model"], rows="clear"
|
|
288
|
+
out["common_model"] = _brain_result(
|
|
289
|
+
sources[0], out["common_model"], "common_model", rows="clear"
|
|
260
290
|
)
|
|
261
291
|
# `_hyperalign` already returns these in the
|
|
262
292
|
# `transformed = original @ T` orientation, and they are square on
|
|
263
293
|
# the voxel axis, so unlike the SRM matrices they are wrapped as-is.
|
|
264
294
|
out["transformation_matrix"] = [
|
|
265
|
-
|
|
295
|
+
_brain_result(source, values, "transformation_matrix", rows="clear")
|
|
266
296
|
for source, values in zip(sources, out["transformation_matrix"])
|
|
267
297
|
]
|
|
268
298
|
else:
|
|
@@ -281,67 +311,28 @@ def align(
|
|
|
281
311
|
# BrainData: (timepoints, voxels)
|
|
282
312
|
# numpy: (voxels, timepoints)
|
|
283
313
|
|
|
284
|
-
|
|
314
|
+
# For procrustes, transformed holds BrainData objects; extract .data.
|
|
315
|
+
# For SRM methods it already holds numpy arrays.
|
|
316
|
+
transformed_arrays = [
|
|
317
|
+
x.data if isinstance(x, BrainData) else x for x in out["transformed"]
|
|
318
|
+
]
|
|
319
|
+
|
|
320
|
+
# Put every case in one orientation, (aligned units, observations), so the
|
|
321
|
+
# correlation below reads the same way whatever came in. BrainData results
|
|
322
|
+
# are (timepoints, voxels) and numpy results are (voxels, timepoints), so
|
|
323
|
+
# exactly one of the two needs a transpose for a given axis.
|
|
324
|
+
if (data_type == "BrainData") == (axis == 0):
|
|
325
|
+
units = [x.T for x in transformed_arrays]
|
|
326
|
+
else:
|
|
327
|
+
units = transformed_arrays
|
|
285
328
|
|
|
286
|
-
|
|
287
|
-
|
|
288
|
-
|
|
289
|
-
|
|
290
|
-
|
|
291
|
-
|
|
292
|
-
]
|
|
293
|
-
if axis == 0:
|
|
294
|
-
# Aligned timepoints → ISC per voxel (correlation over time)
|
|
295
|
-
n_isc = transformed_arrays[0].shape[1] # n_voxels
|
|
296
|
-
for v in range(n_isc):
|
|
297
|
-
# Extract timecourse for voxel v from each subject
|
|
298
|
-
isc_data = np.array([x[:, v] for x in transformed_arrays])
|
|
299
|
-
a = a.append(
|
|
300
|
-
Adjacency(
|
|
301
|
-
1 - pairwise_distances(isc_data, metric="correlation"),
|
|
302
|
-
matrix_type="similarity",
|
|
303
|
-
)
|
|
304
|
-
)
|
|
305
|
-
else: # axis == 1
|
|
306
|
-
# Aligned voxels → ISC per timepoint (spatial correlation)
|
|
307
|
-
n_isc = transformed_arrays[0].shape[0] # n_timepoints
|
|
308
|
-
for t in range(n_isc):
|
|
309
|
-
# Extract spatial pattern at timepoint t from each subject
|
|
310
|
-
isc_data = np.array([x[t, :] for x in transformed_arrays])
|
|
311
|
-
a = a.append(
|
|
312
|
-
Adjacency(
|
|
313
|
-
1 - pairwise_distances(isc_data, metric="correlation"),
|
|
314
|
-
matrix_type="similarity",
|
|
315
|
-
)
|
|
316
|
-
)
|
|
317
|
-
else: # numpy
|
|
318
|
-
# numpy transformed shape: (voxels, timepoints)
|
|
319
|
-
if axis == 0:
|
|
320
|
-
# Aligned timepoints → ISC per voxel (correlation over time)
|
|
321
|
-
n_isc = out["transformed"][0].shape[0] # n_voxels
|
|
322
|
-
for v in range(n_isc):
|
|
323
|
-
# Extract timecourse for voxel v from each subject
|
|
324
|
-
isc_data = np.array([x[v, :] for x in out["transformed"]])
|
|
325
|
-
a = a.append(
|
|
326
|
-
Adjacency(
|
|
327
|
-
1 - pairwise_distances(isc_data, metric="correlation"),
|
|
328
|
-
matrix_type="similarity",
|
|
329
|
-
)
|
|
330
|
-
)
|
|
331
|
-
else: # axis == 1
|
|
332
|
-
# Aligned voxels → ISC per timepoint (spatial correlation)
|
|
333
|
-
n_isc = out["transformed"][0].shape[1] # n_timepoints
|
|
334
|
-
for t in range(n_isc):
|
|
335
|
-
# Extract spatial pattern at timepoint t from each subject
|
|
336
|
-
isc_data = np.array([x[:, t] for x in out["transformed"]])
|
|
337
|
-
a = a.append(
|
|
338
|
-
Adjacency(
|
|
339
|
-
1 - pairwise_distances(isc_data, metric="correlation"),
|
|
340
|
-
matrix_type="similarity",
|
|
341
|
-
)
|
|
342
|
-
)
|
|
343
|
-
|
|
344
|
-
out["isc"] = dict(zip(np.arange(n_isc), a.mean(axis=1)))
|
|
329
|
+
upper_triangle = np.triu_indices(len(units), k=1)
|
|
330
|
+
out["isc"] = {}
|
|
331
|
+
for unit in range(units[0].shape[0]):
|
|
332
|
+
similarity = 1 - pairwise_distances(
|
|
333
|
+
np.array([x[unit] for x in units]), metric="correlation"
|
|
334
|
+
)
|
|
335
|
+
out["isc"][unit] = float(np.nanmean(similarity[upper_triangle]))
|
|
345
336
|
|
|
346
337
|
return out
|
|
347
338
|
|
|
@@ -478,14 +469,12 @@ def procrustes_distance(
|
|
|
478
469
|
# the SAME scale. Previously the observed disparity was compared against a
|
|
479
470
|
# null of similarities, inverting the scales and yielding p ~ 1 for
|
|
480
471
|
# near-identical matrices.
|
|
481
|
-
|
|
482
|
-
observed_similarity = 1 - disparity
|
|
472
|
+
observed_similarity = _procrustes_similarity(mat1, mat2)
|
|
483
473
|
|
|
484
|
-
|
|
485
|
-
delayed(
|
|
474
|
+
null_similarity = Parallel(n_jobs=n_jobs)(
|
|
475
|
+
delayed(_procrustes_similarity)(random_state.permutation(mat1), mat2)
|
|
486
476
|
for _ in range(n_permute)
|
|
487
477
|
)
|
|
488
|
-
null_similarity = [1 - x[2] for x in null_disparities]
|
|
489
478
|
|
|
490
479
|
# Use _compute_pvalue from inference module (signature: obs_stat, null_dist, tail)
|
|
491
480
|
stats = {"similarity": float(observed_similarity)}
|
|
@@ -521,7 +510,8 @@ def align_states(
|
|
|
521
510
|
reordered data. Defaults to False.
|
|
522
511
|
replace_zero_variance (bool): Replace zero-variance columns with uniform
|
|
523
512
|
random numbers before computing distances; avoids NaNs with the
|
|
524
|
-
correlation metric.
|
|
513
|
+
correlation metric. Integer inputs are converted to float so the
|
|
514
|
+
replacement noise survives. Defaults to False.
|
|
525
515
|
|
|
526
516
|
Returns:
|
|
527
517
|
np.ndarray: If `return_index=False` (default), `target[:, remapping]` — the
|
|
@@ -541,12 +531,18 @@ def align_states(
|
|
|
541
531
|
Prevents NaN values when correlation-based distance metrics encounter
|
|
542
532
|
constant columns.
|
|
543
533
|
|
|
534
|
+
The array is converted to float first: writing U(0, 1) draws into an
|
|
535
|
+
integer array truncates every one of them to zero, leaving the constant
|
|
536
|
+
column constant and the correlation distance NaN.
|
|
537
|
+
|
|
544
538
|
Args:
|
|
545
539
|
data (np.ndarray): 2-D array whose columns are checked for zero variance.
|
|
546
540
|
|
|
547
541
|
Returns:
|
|
548
|
-
np.ndarray:
|
|
542
|
+
np.ndarray: Float array with zero-variance columns replaced by
|
|
543
|
+
U(0, 1) values.
|
|
549
544
|
"""
|
|
545
|
+
data = np.asarray(data, dtype=float)
|
|
550
546
|
if np.any(data.std(axis=0) == 0):
|
|
551
547
|
for i in np.where(data.std(axis=0) == 0)[0]:
|
|
552
548
|
data[:, i] = np.random.uniform(low=0, high=1, size=data.shape[0])
|
|
@@ -104,6 +104,62 @@ def _init_w_transforms(
|
|
|
104
104
|
return w, voxels
|
|
105
105
|
|
|
106
106
|
|
|
107
|
+
def _check_n_features(X: list[np.ndarray], n_features: int) -> None:
|
|
108
|
+
"""Reject a feature count no subject's voxel dimension can support.
|
|
109
|
+
|
|
110
|
+
`_init_w_transforms` takes a reduced QR, which silently returns
|
|
111
|
+
`min(voxels, n_features)` columns, so an oversized request would otherwise
|
|
112
|
+
produce a model of a dimension the caller never asked for.
|
|
113
|
+
|
|
114
|
+
Note:
|
|
115
|
+
"Voxels" is this module's name for the first axis throughout, and the
|
|
116
|
+
message follows it. `align(axis=1)` transposes before fitting, so on
|
|
117
|
+
that path the first axis is timepoints and the message still says
|
|
118
|
+
voxels.
|
|
119
|
+
|
|
120
|
+
Args:
|
|
121
|
+
X (list[np.ndarray]): One (voxels_i, samples) array per subject.
|
|
122
|
+
n_features (int): Requested number of shared features.
|
|
123
|
+
|
|
124
|
+
Raises:
|
|
125
|
+
ValueError: If `n_features` is not a positive integer, or exceeds any
|
|
126
|
+
subject's voxel count.
|
|
127
|
+
"""
|
|
128
|
+
if not isinstance(n_features, (int, np.integer)) or n_features < 1:
|
|
129
|
+
raise ValueError(f"n_features must be a positive integer, got {n_features!r}.")
|
|
130
|
+
for subject, data in enumerate(X):
|
|
131
|
+
if data is None:
|
|
132
|
+
continue
|
|
133
|
+
if data.shape[0] < n_features:
|
|
134
|
+
raise ValueError(
|
|
135
|
+
f"subject {subject} has {data.shape[0]} voxels, too few to "
|
|
136
|
+
f"support {n_features} features. Lower n_features to at most "
|
|
137
|
+
"the smallest subject's voxel count."
|
|
138
|
+
)
|
|
139
|
+
|
|
140
|
+
|
|
141
|
+
def _update_transform_subject(Xi: np.ndarray, S: np.ndarray) -> np.ndarray:
|
|
142
|
+
"""Update the mapping $W_i$ for one subject.
|
|
143
|
+
|
|
144
|
+
Solves the orthogonal Procrustes problem
|
|
145
|
+
$\\min ||X_i - W_i S||_F^2$ subject to $W_i^T W_i = I$: with the SVD
|
|
146
|
+
$U \\Sigma V^T = X_i S^T$, the optimum is $W_i = U V^T$.
|
|
147
|
+
|
|
148
|
+
Args:
|
|
149
|
+
Xi (np.ndarray): The subject's data $X_i$, shape (voxels, timepoints).
|
|
150
|
+
S (np.ndarray): The shared response, shape (n_features, timepoints).
|
|
151
|
+
|
|
152
|
+
Returns:
|
|
153
|
+
np.ndarray: The orthogonal transform $W_i$, shape (voxels, n_features).
|
|
154
|
+
"""
|
|
155
|
+
# Compute cross-covariance: X_i S^T
|
|
156
|
+
A = Xi.dot(S.T)
|
|
157
|
+
# Solve the Procrustes problem via SVD
|
|
158
|
+
# Optimal orthogonal transform: W_i = U V^T where A = U Σ V^T
|
|
159
|
+
U, _, V = np.linalg.svd(A, full_matrices=False)
|
|
160
|
+
return U.dot(V)
|
|
161
|
+
|
|
162
|
+
|
|
107
163
|
class _SRM(BaseEstimator, TransformerMixin):
|
|
108
164
|
"""Probabilistic Shared Response Model (SRM).
|
|
109
165
|
|
|
@@ -142,7 +198,7 @@ class _SRM(BaseEstimator, TransformerMixin):
|
|
|
142
198
|
Examples:
|
|
143
199
|
```python
|
|
144
200
|
import numpy as np
|
|
145
|
-
from nltools.algorithms import
|
|
201
|
+
from nltools.algorithms.alignment import _SRM
|
|
146
202
|
|
|
147
203
|
data = [np.random.randn(100, 50) for _ in range(3)] # 3 subjects
|
|
148
204
|
|
|
@@ -173,6 +229,11 @@ class _SRM(BaseEstimator, TransformerMixin):
|
|
|
173
229
|
|
|
174
230
|
Returns:
|
|
175
231
|
_SRM: Fitted model (`self`).
|
|
232
|
+
|
|
233
|
+
Raises:
|
|
234
|
+
ValueError: If there are fewer than two subjects, the subjects
|
|
235
|
+
disagree on sample count, or `n_features` exceeds any subject's
|
|
236
|
+
voxel count.
|
|
176
237
|
"""
|
|
177
238
|
logger.info("Starting Probabilistic SRM")
|
|
178
239
|
|
|
@@ -197,6 +258,10 @@ class _SRM(BaseEstimator, TransformerMixin):
|
|
|
197
258
|
f"Different number of samples between subjects: {sample_counts}."
|
|
198
259
|
)
|
|
199
260
|
|
|
261
|
+
# After the sample-count check, so input wrong on both axes reports the
|
|
262
|
+
# mismatched samples first.
|
|
263
|
+
_check_n_features(X, self.n_features)
|
|
264
|
+
|
|
200
265
|
# Validate all data is finite
|
|
201
266
|
for subject in range(number_subjects):
|
|
202
267
|
if X[subject] is not None:
|
|
@@ -328,28 +393,6 @@ class _SRM(BaseEstimator, TransformerMixin):
|
|
|
328
393
|
|
|
329
394
|
return loglikehood
|
|
330
395
|
|
|
331
|
-
@staticmethod
|
|
332
|
-
def _update_transform_subject(Xi, S):
|
|
333
|
-
"""Update the mapping $W_i$ for one subject.
|
|
334
|
-
|
|
335
|
-
Solves the orthogonal Procrustes problem
|
|
336
|
-
$\\min ||X_i - W_i S||_F^2$ subject to $W_i^T W_i = I$: with the SVD
|
|
337
|
-
$U \\Sigma V^T = X_i S^T$, the optimum is $W_i = U V^T$.
|
|
338
|
-
|
|
339
|
-
Args:
|
|
340
|
-
Xi (np.ndarray): The subject's data $X_i$, shape (voxels, timepoints).
|
|
341
|
-
S (np.ndarray): The shared response, shape (n_features, timepoints).
|
|
342
|
-
|
|
343
|
-
Returns:
|
|
344
|
-
np.ndarray: The orthogonal transform $W_i$, shape (voxels, n_features).
|
|
345
|
-
"""
|
|
346
|
-
# Compute cross-covariance: X_i S^T
|
|
347
|
-
A = Xi.dot(S.T)
|
|
348
|
-
# Solve the Procrustes problem via SVD
|
|
349
|
-
# Optimal orthogonal transform: W_i = U V^T where A = U Σ V^T
|
|
350
|
-
U, _, V = np.linalg.svd(A, full_matrices=False)
|
|
351
|
-
return U.dot(V)
|
|
352
|
-
|
|
353
396
|
def transform_subject(self, X: np.ndarray) -> np.ndarray:
|
|
354
397
|
"""Transform a new subject using the existing model.
|
|
355
398
|
|
|
@@ -373,7 +416,7 @@ class _SRM(BaseEstimator, TransformerMixin):
|
|
|
373
416
|
"The number of timepoints(TRs) does not match the one in the model."
|
|
374
417
|
)
|
|
375
418
|
|
|
376
|
-
w =
|
|
419
|
+
w = _update_transform_subject(X, self.s_)
|
|
377
420
|
|
|
378
421
|
return w
|
|
379
422
|
|
|
@@ -533,7 +576,7 @@ class _DetSRM(BaseEstimator, TransformerMixin):
|
|
|
533
576
|
Examples:
|
|
534
577
|
```python
|
|
535
578
|
import numpy as np
|
|
536
|
-
from nltools.algorithms import
|
|
579
|
+
from nltools.algorithms.alignment import _DetSRM
|
|
537
580
|
|
|
538
581
|
data = [np.random.randn(100, 50) for _ in range(3)] # 3 subjects
|
|
539
582
|
|
|
@@ -563,6 +606,11 @@ class _DetSRM(BaseEstimator, TransformerMixin):
|
|
|
563
606
|
|
|
564
607
|
Returns:
|
|
565
608
|
_DetSRM: Fitted model (`self`).
|
|
609
|
+
|
|
610
|
+
Raises:
|
|
611
|
+
ValueError: If there are fewer than two subjects, the subjects
|
|
612
|
+
disagree on sample count, or `n_features` exceeds any subject's
|
|
613
|
+
voxel count.
|
|
566
614
|
"""
|
|
567
615
|
logger.info("Starting Deterministic SRM")
|
|
568
616
|
|
|
@@ -587,6 +635,10 @@ class _DetSRM(BaseEstimator, TransformerMixin):
|
|
|
587
635
|
if X[subject].shape[1] != number_trs:
|
|
588
636
|
raise ValueError("Different number of samples between subjects.")
|
|
589
637
|
|
|
638
|
+
# After the sample-count check, so input wrong on both axes reports the
|
|
639
|
+
# mismatched samples first.
|
|
640
|
+
_check_n_features(X, self.n_features)
|
|
641
|
+
|
|
590
642
|
# Run SRM
|
|
591
643
|
self.w_, self.s_ = self._srm(X)
|
|
592
644
|
|
|
@@ -655,28 +707,6 @@ class _DetSRM(BaseEstimator, TransformerMixin):
|
|
|
655
707
|
|
|
656
708
|
return s
|
|
657
709
|
|
|
658
|
-
@staticmethod
|
|
659
|
-
def _update_transform_subject(Xi, S):
|
|
660
|
-
"""Update the mapping $W_i$ for one subject.
|
|
661
|
-
|
|
662
|
-
Solves the orthogonal Procrustes problem
|
|
663
|
-
$\\min ||X_i - W_i S||_F^2$ subject to $W_i^T W_i = I$: with the SVD
|
|
664
|
-
$U \\Sigma V^T = X_i S^T$, the optimum is $W_i = U V^T$.
|
|
665
|
-
|
|
666
|
-
Args:
|
|
667
|
-
Xi (np.ndarray): The subject's data $X_i$, shape (voxels, timepoints).
|
|
668
|
-
S (np.ndarray): The shared response, shape (n_features, timepoints).
|
|
669
|
-
|
|
670
|
-
Returns:
|
|
671
|
-
np.ndarray: The orthogonal transform $W_i$, shape (voxels, n_features).
|
|
672
|
-
"""
|
|
673
|
-
# Compute cross-covariance: X_i S^T
|
|
674
|
-
A = Xi.dot(S.T)
|
|
675
|
-
# Solve the Procrustes problem via SVD
|
|
676
|
-
# Optimal orthogonal transform: W_i = U V^T where A = U Σ V^T
|
|
677
|
-
U, _, V = np.linalg.svd(A, full_matrices=False)
|
|
678
|
-
return U.dot(V)
|
|
679
|
-
|
|
680
710
|
def transform_subject(self, X: np.ndarray) -> np.ndarray:
|
|
681
711
|
"""Transform a new subject using the existing model.
|
|
682
712
|
|
|
@@ -700,7 +730,7 @@ class _DetSRM(BaseEstimator, TransformerMixin):
|
|
|
700
730
|
"The number of timepoints(TRs) does not match the one in the model."
|
|
701
731
|
)
|
|
702
732
|
|
|
703
|
-
w =
|
|
733
|
+
w = _update_transform_subject(X, self.s_)
|
|
704
734
|
|
|
705
735
|
return w
|
|
706
736
|
|
|
@@ -654,6 +654,10 @@ def _ridge_bootstrap_batch_size(
|
|
|
654
654
|
#: every summary payload is converted to CPU float64 before it is retained.
|
|
655
655
|
_BOOTSTRAP_OUTPUT_ITEMSIZE = 8
|
|
656
656
|
|
|
657
|
+
#: Bytes per resampling index. `_generate_bootstrap_indices` returns an int64
|
|
658
|
+
#: matrix, and every CPU worker's closure captures it.
|
|
659
|
+
_BOOTSTRAP_INDEX_ITEMSIZE = 8
|
|
660
|
+
|
|
657
661
|
#: Output-sized arrays a bootstrap run always holds beyond its retained
|
|
658
662
|
#: replicates: the two Welford accumulators (running mean and running sum of
|
|
659
663
|
#: squared deviations) and the four `BootstrapResult` summary payloads.
|
|
@@ -734,15 +738,19 @@ def _bootstrap_output_bytes(
|
|
|
734
738
|
*,
|
|
735
739
|
confidence_level: float,
|
|
736
740
|
return_samples: bool,
|
|
741
|
+
n_obs: int,
|
|
737
742
|
n_workers: int = 1,
|
|
738
743
|
) -> int:
|
|
739
|
-
"""Bytes a bootstrap run must hold for its retained output.
|
|
744
|
+
"""Bytes a bootstrap run must hold for its retained output and its indices.
|
|
740
745
|
|
|
741
746
|
Charges eight bytes for every output-sized array a run holds at once: the
|
|
742
747
|
two bounded tails, the replicates buffered before the next flush and the
|
|
743
748
|
two temporaries that flush builds, one dispatch window of in-flight
|
|
744
749
|
replicates, every replicate when `return_samples=True`, and the two Welford
|
|
745
|
-
accumulators plus the four summary payloads.
|
|
750
|
+
accumulators plus the four summary payloads. On top of that it charges the
|
|
751
|
+
resampling index matrix twice: the run retains one `(n_samples, n_obs)`
|
|
752
|
+
int64 array, and building it holds the per-draw vectors alongside the
|
|
753
|
+
stacked result.
|
|
746
754
|
|
|
747
755
|
Args:
|
|
748
756
|
output_shape (tuple[int, ...]): Shape of one replicate's output.
|
|
@@ -750,6 +758,8 @@ def _bootstrap_output_bytes(
|
|
|
750
758
|
confidence_level (float): Interval confidence level, which sets the
|
|
751
759
|
retained tail size.
|
|
752
760
|
return_samples (bool): Whether the complete distribution is retained.
|
|
761
|
+
n_obs (int): Observations resampled per replicate, which sizes the
|
|
762
|
+
index matrix.
|
|
753
763
|
n_workers (int): Planned CPU worker count, which sets the dispatch
|
|
754
764
|
window. Defaults to 1 (the GPU driver budgets its own batch through
|
|
755
765
|
`_ridge_bootstrap_batch_size` instead).
|
|
@@ -770,7 +780,8 @@ def _bootstrap_output_bytes(
|
|
|
770
780
|
+ (int(n_samples) if return_samples else 0)
|
|
771
781
|
+ _BOOTSTRAP_FIXED_OUTPUT_ARRAYS
|
|
772
782
|
)
|
|
773
|
-
|
|
783
|
+
index_bytes = 2 * int(n_obs) * int(n_samples) * _BOOTSTRAP_INDEX_ITEMSIZE
|
|
784
|
+
return output_size * arrays * _BOOTSTRAP_OUTPUT_ITEMSIZE + index_bytes
|
|
774
785
|
|
|
775
786
|
|
|
776
787
|
def _bootstrap_memory_preflight(
|
|
@@ -779,6 +790,7 @@ def _bootstrap_memory_preflight(
|
|
|
779
790
|
*,
|
|
780
791
|
confidence_level: float,
|
|
781
792
|
return_samples: bool,
|
|
793
|
+
n_obs: int,
|
|
782
794
|
n_workers: int = 1,
|
|
783
795
|
memory_budget_gb: float | None = None,
|
|
784
796
|
backend: "_Backend | None" = None,
|
|
@@ -795,6 +807,8 @@ def _bootstrap_memory_preflight(
|
|
|
795
807
|
n_samples (int): Number of bootstrap replicates.
|
|
796
808
|
confidence_level (float): Interval confidence level.
|
|
797
809
|
return_samples (bool): Whether the complete distribution is retained.
|
|
810
|
+
n_obs (int): Observations resampled per replicate, which sizes the
|
|
811
|
+
index matrix.
|
|
798
812
|
n_workers (int): Planned CPU worker count, which sets the dispatch
|
|
799
813
|
window. Defaults to 1.
|
|
800
814
|
memory_budget_gb (float | None): Explicit budget in GB, or None to
|
|
@@ -813,6 +827,7 @@ def _bootstrap_memory_preflight(
|
|
|
813
827
|
n_samples,
|
|
814
828
|
confidence_level=confidence_level,
|
|
815
829
|
return_samples=return_samples,
|
|
830
|
+
n_obs=n_obs,
|
|
816
831
|
n_workers=n_workers,
|
|
817
832
|
)
|
|
818
833
|
required_gb = required_bytes / 1e9
|
|
@@ -925,7 +940,9 @@ def _compute_oom_safe(fn, *arrays, min_chunk: int = 1):
|
|
|
925
940
|
a numpy array whose axis 0 corresponds row-for-row to its inputs. On a
|
|
926
941
|
device OOM the cache is emptied, the arrays are split in half along
|
|
927
942
|
axis 0, and the halves are retried recursively; partial results are
|
|
928
|
-
concatenated along axis 0.
|
|
943
|
+
concatenated along axis 0. Recovery happens after leaving the exception
|
|
944
|
+
handler, so the failed call's allocations are released before the smaller
|
|
945
|
+
retries ask for them.
|
|
929
946
|
|
|
930
947
|
Because splitting reuses the *already generated* inputs rather than
|
|
931
948
|
re-drawing them, recovery never changes which permutations a seeded
|
|
@@ -954,17 +971,22 @@ def _compute_oom_safe(fn, *arrays, min_chunk: int = 1):
|
|
|
954
971
|
except Exception as exc:
|
|
955
972
|
if not _is_oom_error(exc):
|
|
956
973
|
raise
|
|
957
|
-
_empty_device_cache()
|
|
958
974
|
if n <= min_chunk:
|
|
975
|
+
_empty_device_cache()
|
|
959
976
|
raise MemoryError(
|
|
960
977
|
f"Device out of memory even for a single item (chunk of {n}). "
|
|
961
978
|
"Reduce the problem size, lower max_gpu_memory_gb elsewhere on "
|
|
962
979
|
"the device, or use device='cpu'."
|
|
963
980
|
) from exc
|
|
964
|
-
|
|
965
|
-
|
|
966
|
-
|
|
967
|
-
|
|
981
|
+
|
|
982
|
+
# Outside the handler, where Python's implicit `del exc` has dropped the
|
|
983
|
+
# traceback: the failed call's frame — and the device allocations its
|
|
984
|
+
# locals held — are gone before the cache is emptied and the halves retried.
|
|
985
|
+
_empty_device_cache()
|
|
986
|
+
mid = n // 2
|
|
987
|
+
left = _compute_oom_safe(fn, *(a[:mid] for a in arrays), min_chunk=min_chunk)
|
|
988
|
+
right = _compute_oom_safe(fn, *(a[mid:] for a in arrays), min_chunk=min_chunk)
|
|
989
|
+
return np.concatenate([left, right], axis=0)
|
|
968
990
|
|
|
969
991
|
|
|
970
992
|
# ----------------------------------------------------------------------
|