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,282 @@
|
|
|
1
|
+
"""Shared input validation for the inference module.
|
|
2
|
+
|
|
3
|
+
One home for the argument checks the permutation, bootstrap, and matrix tests
|
|
4
|
+
share, so every entry point raises the same `ValueError` for the same mistake.
|
|
5
|
+
The `tail` vocabulary is wider than inference, so it lives one level up in
|
|
6
|
+
`nltools.algorithms.validation`.
|
|
7
|
+
|
|
8
|
+
Examples:
|
|
9
|
+
```python
|
|
10
|
+
from nltools.algorithms.inference.validation import _validate_square_matrix
|
|
11
|
+
|
|
12
|
+
_validate_square_matrix(np.eye(3)) # → None
|
|
13
|
+
_validate_square_matrix(np.zeros((2, 3))) # raises ValueError
|
|
14
|
+
```
|
|
15
|
+
"""
|
|
16
|
+
|
|
17
|
+
import numpy as np
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
def _validate_array_shape(
|
|
21
|
+
array: np.ndarray,
|
|
22
|
+
expected_ndim: int,
|
|
23
|
+
name: str = "array",
|
|
24
|
+
) -> None:
|
|
25
|
+
"""Validate array dimensionality.
|
|
26
|
+
|
|
27
|
+
Args:
|
|
28
|
+
array (np.ndarray): Array to validate.
|
|
29
|
+
expected_ndim (int): Expected number of dimensions.
|
|
30
|
+
name (str): Name of the array for the error message.
|
|
31
|
+
|
|
32
|
+
Raises:
|
|
33
|
+
ValueError: If the array has the wrong number of dimensions.
|
|
34
|
+
"""
|
|
35
|
+
if array.ndim != expected_ndim:
|
|
36
|
+
raise ValueError(
|
|
37
|
+
f"{name} must be {expected_ndim}D, got shape {array.shape} ({array.ndim}D)"
|
|
38
|
+
)
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
def _validate_array_shape_range(
|
|
42
|
+
array: np.ndarray,
|
|
43
|
+
min_ndim: int,
|
|
44
|
+
max_ndim: int,
|
|
45
|
+
name: str = "array",
|
|
46
|
+
) -> None:
|
|
47
|
+
"""Validate that array dimensionality falls within a range.
|
|
48
|
+
|
|
49
|
+
Args:
|
|
50
|
+
array (np.ndarray): Array to validate.
|
|
51
|
+
min_ndim (int): Minimum number of dimensions (inclusive).
|
|
52
|
+
max_ndim (int): Maximum number of dimensions (inclusive).
|
|
53
|
+
name (str): Name of the array for the error message.
|
|
54
|
+
|
|
55
|
+
Raises:
|
|
56
|
+
ValueError: If the array has the wrong number of dimensions.
|
|
57
|
+
"""
|
|
58
|
+
if not (min_ndim <= array.ndim <= max_ndim):
|
|
59
|
+
raise ValueError(
|
|
60
|
+
f"{name} must be {min_ndim}D to {max_ndim}D, got shape {array.shape} "
|
|
61
|
+
f"({array.ndim}D)"
|
|
62
|
+
)
|
|
63
|
+
|
|
64
|
+
|
|
65
|
+
def _validate_same_shape(
|
|
66
|
+
array1: np.ndarray,
|
|
67
|
+
array2: np.ndarray,
|
|
68
|
+
name1: str = "array1",
|
|
69
|
+
name2: str = "array2",
|
|
70
|
+
) -> None:
|
|
71
|
+
"""Validate that two arrays have the same shape.
|
|
72
|
+
|
|
73
|
+
Args:
|
|
74
|
+
array1 (np.ndarray): First array.
|
|
75
|
+
array2 (np.ndarray): Second array.
|
|
76
|
+
name1 (str): Name of the first array for the error message.
|
|
77
|
+
name2 (str): Name of the second array for the error message.
|
|
78
|
+
|
|
79
|
+
Raises:
|
|
80
|
+
ValueError: If the arrays have different shapes.
|
|
81
|
+
"""
|
|
82
|
+
if array1.shape != array2.shape:
|
|
83
|
+
raise ValueError(
|
|
84
|
+
f"{name1} and {name2} must have same shape, "
|
|
85
|
+
f"got {array1.shape} and {array2.shape}"
|
|
86
|
+
)
|
|
87
|
+
|
|
88
|
+
|
|
89
|
+
def _validate_metric_parameter(
|
|
90
|
+
metric: str,
|
|
91
|
+
allowed: list[str],
|
|
92
|
+
name: str = "metric",
|
|
93
|
+
) -> None:
|
|
94
|
+
"""Validate a metric name against an allowed list.
|
|
95
|
+
|
|
96
|
+
Args:
|
|
97
|
+
metric (str): Metric name to validate.
|
|
98
|
+
allowed (list[str]): Allowed metric names.
|
|
99
|
+
name (str): Name of the parameter for the error message.
|
|
100
|
+
|
|
101
|
+
Raises:
|
|
102
|
+
ValueError: If `metric` is not in `allowed`.
|
|
103
|
+
"""
|
|
104
|
+
if metric not in allowed:
|
|
105
|
+
allowed_str = ", ".join(f"'{m}'" for m in allowed)
|
|
106
|
+
raise ValueError(f"{name} must be one of [{allowed_str}], got {metric!r}")
|
|
107
|
+
|
|
108
|
+
|
|
109
|
+
def _validate_how_parameter(how: str) -> None:
|
|
110
|
+
"""Validate the `how` parameter for matrix operations.
|
|
111
|
+
|
|
112
|
+
Args:
|
|
113
|
+
how (str): `'upper'`, `'lower'`, or `'full'`.
|
|
114
|
+
|
|
115
|
+
Raises:
|
|
116
|
+
ValueError: If `how` is not one of those values.
|
|
117
|
+
"""
|
|
118
|
+
if how not in ["upper", "lower", "full"]:
|
|
119
|
+
raise ValueError(f"how must be 'upper', 'lower', or 'full', got {how!r}")
|
|
120
|
+
|
|
121
|
+
|
|
122
|
+
def _validate_square_matrix(matrix: np.ndarray, name: str = "matrix") -> None:
|
|
123
|
+
"""Validate that a matrix is square.
|
|
124
|
+
|
|
125
|
+
Args:
|
|
126
|
+
matrix (np.ndarray): Matrix to validate.
|
|
127
|
+
name (str): Name of the matrix for the error message.
|
|
128
|
+
|
|
129
|
+
Raises:
|
|
130
|
+
ValueError: If the matrix is not square.
|
|
131
|
+
"""
|
|
132
|
+
if matrix.shape[0] != matrix.shape[1]:
|
|
133
|
+
raise ValueError(f"{name} must be square, got shape {matrix.shape}")
|
|
134
|
+
|
|
135
|
+
|
|
136
|
+
def _validate_shape_compatibility(
|
|
137
|
+
X: np.ndarray,
|
|
138
|
+
y: np.ndarray,
|
|
139
|
+
X_name: str = "X",
|
|
140
|
+
y_name: str = "y",
|
|
141
|
+
) -> None:
|
|
142
|
+
"""Validate that X and y have the same number of samples.
|
|
143
|
+
|
|
144
|
+
Args:
|
|
145
|
+
X (np.ndarray): Feature matrix.
|
|
146
|
+
y (np.ndarray): Target vector or matrix.
|
|
147
|
+
X_name (str): Name of X for the error message.
|
|
148
|
+
y_name (str): Name of y for the error message.
|
|
149
|
+
|
|
150
|
+
Raises:
|
|
151
|
+
ValueError: If the first dimensions differ.
|
|
152
|
+
"""
|
|
153
|
+
if X.shape[0] != y.shape[0]:
|
|
154
|
+
raise ValueError(
|
|
155
|
+
f"{X_name} and {y_name} must have same first dimension (n_samples), "
|
|
156
|
+
f"got {X.shape[0]} and {y.shape[0]}"
|
|
157
|
+
)
|
|
158
|
+
|
|
159
|
+
|
|
160
|
+
def _validate_bootstrap_method(
|
|
161
|
+
method: str, simple_methods: list[str], fitted_methods: list[str]
|
|
162
|
+
) -> None:
|
|
163
|
+
"""Validate a bootstrap method name.
|
|
164
|
+
|
|
165
|
+
Args:
|
|
166
|
+
method (str): Method name to validate.
|
|
167
|
+
simple_methods (list[str]): Methods that need no fitted model.
|
|
168
|
+
fitted_methods (list[str]): Methods that require a prior `.fit()`.
|
|
169
|
+
|
|
170
|
+
Raises:
|
|
171
|
+
ValueError: If `method` is in neither list.
|
|
172
|
+
"""
|
|
173
|
+
supported = simple_methods + fitted_methods
|
|
174
|
+
if method not in supported:
|
|
175
|
+
raise ValueError(
|
|
176
|
+
f"Unsupported method '{method}'. "
|
|
177
|
+
f"Supported methods: {simple_methods} (simple methods), "
|
|
178
|
+
f"{fitted_methods} (fitted model methods). "
|
|
179
|
+
f"For fitted methods, you must call .fit() first."
|
|
180
|
+
)
|
|
181
|
+
|
|
182
|
+
|
|
183
|
+
def _validate_bootstrap_data(data: np.ndarray, method: str) -> None:
|
|
184
|
+
"""Validate input data for bootstrapping.
|
|
185
|
+
|
|
186
|
+
Args:
|
|
187
|
+
data (np.ndarray): 1D or 2D data with at least 2 samples along axis 0.
|
|
188
|
+
method (str): Bootstrap method name (reserved for method-specific checks).
|
|
189
|
+
|
|
190
|
+
Raises:
|
|
191
|
+
ValueError: If the data is not 1D/2D or has fewer than 2 samples.
|
|
192
|
+
"""
|
|
193
|
+
# Check dimensionality
|
|
194
|
+
if data.ndim not in [1, 2]:
|
|
195
|
+
raise ValueError(
|
|
196
|
+
f"Data must be 1D or 2D, got shape {data.shape}. "
|
|
197
|
+
f"For 3D+ data, you may need to reshape or select specific dimensions."
|
|
198
|
+
)
|
|
199
|
+
|
|
200
|
+
# Check number of samples
|
|
201
|
+
n_samples = data.shape[0] if data.ndim == 2 else len(data)
|
|
202
|
+
if n_samples < 2:
|
|
203
|
+
raise ValueError(
|
|
204
|
+
f"Need at least 2 samples for bootstrap, got {n_samples}. "
|
|
205
|
+
f"Bootstrap requires resampling, which needs multiple samples."
|
|
206
|
+
)
|
|
207
|
+
|
|
208
|
+
|
|
209
|
+
def _validate_n_samples(n_samples: int) -> None:
|
|
210
|
+
"""Reject a replicate count a bootstrap cannot be computed from.
|
|
211
|
+
|
|
212
|
+
Two replicates are the fewest a `ddof=1` standard error can be computed
|
|
213
|
+
from, so that is the hard floor. The separate quality advisory lives in
|
|
214
|
+
`_advise_on_n_samples`, in `nltools/algorithms/inference/bootstrap.py`.
|
|
215
|
+
|
|
216
|
+
Args:
|
|
217
|
+
n_samples (int): Number of bootstrap replicates.
|
|
218
|
+
|
|
219
|
+
Raises:
|
|
220
|
+
TypeError: If `n_samples` is not an integer.
|
|
221
|
+
ValueError: If `n_samples` is below 2.
|
|
222
|
+
"""
|
|
223
|
+
if isinstance(n_samples, bool) or not isinstance(n_samples, (int, np.integer)):
|
|
224
|
+
raise TypeError(f"n_samples must be an integer, got {type(n_samples).__name__}")
|
|
225
|
+
|
|
226
|
+
if n_samples < 2:
|
|
227
|
+
raise ValueError(
|
|
228
|
+
f"n_samples must be at least 2, got {n_samples}. "
|
|
229
|
+
f"A bootstrap standard error needs at least two replicates. "
|
|
230
|
+
f"Recommended: n_samples >= 1000 for confidence intervals."
|
|
231
|
+
)
|
|
232
|
+
|
|
233
|
+
|
|
234
|
+
def _validate_confidence_level(confidence_level: float) -> None:
|
|
235
|
+
"""Validate the interval confidence level.
|
|
236
|
+
|
|
237
|
+
Args:
|
|
238
|
+
confidence_level (float): Requested level.
|
|
239
|
+
|
|
240
|
+
Raises:
|
|
241
|
+
TypeError: If `confidence_level` is not a real number.
|
|
242
|
+
ValueError: If it is not finite and strictly between zero and one.
|
|
243
|
+
"""
|
|
244
|
+
if isinstance(confidence_level, bool) or not isinstance(
|
|
245
|
+
confidence_level, (int, float, np.integer, np.floating)
|
|
246
|
+
):
|
|
247
|
+
raise TypeError(
|
|
248
|
+
f"confidence_level must be a number, got {type(confidence_level).__name__}"
|
|
249
|
+
)
|
|
250
|
+
value = float(confidence_level)
|
|
251
|
+
if not np.isfinite(value) or not 0 < value < 1:
|
|
252
|
+
raise ValueError(
|
|
253
|
+
f"confidence_level must be finite and strictly between 0 and 1, got "
|
|
254
|
+
f"{confidence_level!r}. Use 0.95 for a 95% interval."
|
|
255
|
+
)
|
|
256
|
+
|
|
257
|
+
|
|
258
|
+
def _validate_memory_budget(memory_budget_gb: float | None) -> None:
|
|
259
|
+
"""Validate an explicit working-memory budget.
|
|
260
|
+
|
|
261
|
+
Args:
|
|
262
|
+
memory_budget_gb (float | None): Budget in GB, or None to measure the
|
|
263
|
+
device.
|
|
264
|
+
|
|
265
|
+
Raises:
|
|
266
|
+
TypeError: If a supplied budget is not a real number.
|
|
267
|
+
ValueError: If a supplied budget is not finite and positive.
|
|
268
|
+
"""
|
|
269
|
+
if memory_budget_gb is None:
|
|
270
|
+
return
|
|
271
|
+
if isinstance(memory_budget_gb, bool) or not isinstance(
|
|
272
|
+
memory_budget_gb, (int, float, np.integer, np.floating)
|
|
273
|
+
):
|
|
274
|
+
raise TypeError(
|
|
275
|
+
f"memory_budget_gb must be a number or None, got "
|
|
276
|
+
f"{type(memory_budget_gb).__name__}"
|
|
277
|
+
)
|
|
278
|
+
value = float(memory_budget_gb)
|
|
279
|
+
if not np.isfinite(value) or value <= 0:
|
|
280
|
+
raise ValueError(
|
|
281
|
+
f"memory_budget_gb must be finite and positive, got {memory_budget_gb!r}."
|
|
282
|
+
)
|
|
@@ -0,0 +1,207 @@
|
|
|
1
|
+
"""Spatial neighborhood computation for neuroimaging analyses.
|
|
2
|
+
|
|
3
|
+
This module computes spatial neighborhoods (spheres) around brain voxels. It is
|
|
4
|
+
designed to support searchlight analyses, ISC, and other operations that require
|
|
5
|
+
iterating over local brain regions.
|
|
6
|
+
|
|
7
|
+
Examples:
|
|
8
|
+
```python
|
|
9
|
+
import nibabel as nib
|
|
10
|
+
from nltools.algorithms.neighborhoods import compute_searchlight_neighborhoods
|
|
11
|
+
|
|
12
|
+
mask = nib.load("mask.nii.gz")
|
|
13
|
+
neighborhoods = compute_searchlight_neighborhoods(mask, radius=10.0)
|
|
14
|
+
|
|
15
|
+
# Iterate over all voxels and their neighborhoods
|
|
16
|
+
for center_idx, neighbor_indices in neighborhoods.iter_neighborhoods():
|
|
17
|
+
local_data = data[:, neighbor_indices] # data for the voxels in this sphere
|
|
18
|
+
result[center_idx] = analyze(local_data)
|
|
19
|
+
```
|
|
20
|
+
"""
|
|
21
|
+
|
|
22
|
+
from __future__ import annotations
|
|
23
|
+
|
|
24
|
+
from dataclasses import dataclass
|
|
25
|
+
from typing import TYPE_CHECKING
|
|
26
|
+
from collections.abc import Iterator
|
|
27
|
+
|
|
28
|
+
import numpy as np
|
|
29
|
+
from scipy import sparse
|
|
30
|
+
from sklearn import neighbors
|
|
31
|
+
|
|
32
|
+
from nltools.utils import _maybe_tqdm
|
|
33
|
+
|
|
34
|
+
if TYPE_CHECKING:
|
|
35
|
+
from nibabel import Nifti1Image
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
@dataclass(frozen=True)
|
|
39
|
+
class _SphereNeighborhoods:
|
|
40
|
+
"""Precomputed sphere neighborhoods for a brain mask.
|
|
41
|
+
|
|
42
|
+
This dataclass stores a sparse adjacency matrix where row i contains True
|
|
43
|
+
for all voxels within the specified radius of voxel i. It provides efficient
|
|
44
|
+
iteration over neighborhoods for searchlight-style analyses.
|
|
45
|
+
|
|
46
|
+
Attributes:
|
|
47
|
+
adjacency (sparse.csr_matrix): ``(n_voxels, n_voxels)`` matrix where
|
|
48
|
+
``adjacency[i, j]`` is nonzero if voxel ``j`` is within the radius
|
|
49
|
+
of voxel ``i``.
|
|
50
|
+
radius (float): Radius in millimeters.
|
|
51
|
+
n_voxels (int): Number of voxels in the mask.
|
|
52
|
+
mean_size (float): Mean neighborhood size in voxels.
|
|
53
|
+
min_size (int): Smallest neighborhood size in voxels.
|
|
54
|
+
max_size (int): Largest neighborhood size in voxels.
|
|
55
|
+
|
|
56
|
+
Examples:
|
|
57
|
+
```python
|
|
58
|
+
neighborhoods = compute_searchlight_neighborhoods(mask, radius=10.0)
|
|
59
|
+
print(f"Mean neighborhood size: {neighborhoods.mean_size:.1f} voxels")
|
|
60
|
+
|
|
61
|
+
# Get neighbors of a specific voxel
|
|
62
|
+
neighbor_idx = neighborhoods.get_neighbors(100)
|
|
63
|
+
print(f"Voxel 100 has {len(neighbor_idx)} neighbors")
|
|
64
|
+
```
|
|
65
|
+
"""
|
|
66
|
+
|
|
67
|
+
adjacency: sparse.csr_matrix
|
|
68
|
+
radius: float
|
|
69
|
+
n_voxels: int
|
|
70
|
+
|
|
71
|
+
def get_neighbors(self, voxel_idx: int) -> np.ndarray:
|
|
72
|
+
"""Get indices of all voxels in the neighborhood of a given voxel.
|
|
73
|
+
|
|
74
|
+
Args:
|
|
75
|
+
voxel_idx: Index of the center voxel (0 to n_voxels-1)
|
|
76
|
+
|
|
77
|
+
Returns:
|
|
78
|
+
Array of voxel indices within radius of the center voxel
|
|
79
|
+
"""
|
|
80
|
+
return self.adjacency[voxel_idx].indices
|
|
81
|
+
|
|
82
|
+
def get_neighborhood_size(self, voxel_idx: int) -> int:
|
|
83
|
+
"""Get the number of voxels in a neighborhood.
|
|
84
|
+
|
|
85
|
+
Args:
|
|
86
|
+
voxel_idx: Index of the center voxel
|
|
87
|
+
|
|
88
|
+
Returns:
|
|
89
|
+
Number of voxels in the neighborhood
|
|
90
|
+
"""
|
|
91
|
+
return self.adjacency[voxel_idx].nnz
|
|
92
|
+
|
|
93
|
+
def iter_neighborhoods(
|
|
94
|
+
self, *, progress_bar: bool = False
|
|
95
|
+
) -> Iterator[tuple[int, np.ndarray]]:
|
|
96
|
+
"""Iterate over all neighborhoods.
|
|
97
|
+
|
|
98
|
+
Args:
|
|
99
|
+
progress_bar: If True, wrap the iterator with a tqdm progress bar.
|
|
100
|
+
|
|
101
|
+
Yields:
|
|
102
|
+
tuple[int, np.ndarray]: ``(center_voxel_idx, neighbor_indices)`` for
|
|
103
|
+
each voxel.
|
|
104
|
+
"""
|
|
105
|
+
iterator = _maybe_tqdm(
|
|
106
|
+
range(self.n_voxels),
|
|
107
|
+
progress_bar=progress_bar,
|
|
108
|
+
desc="Searchlight",
|
|
109
|
+
unit="voxels",
|
|
110
|
+
)
|
|
111
|
+
|
|
112
|
+
for i in iterator:
|
|
113
|
+
yield i, self.get_neighbors(i)
|
|
114
|
+
|
|
115
|
+
@property
|
|
116
|
+
def mean_size(self) -> float:
|
|
117
|
+
"""Mean neighborhood size in voxels."""
|
|
118
|
+
return float(self.adjacency.sum() / self.n_voxels)
|
|
119
|
+
|
|
120
|
+
@property
|
|
121
|
+
def min_size(self) -> int:
|
|
122
|
+
"""Minimum neighborhood size."""
|
|
123
|
+
sizes = np.diff(self.adjacency.indptr)
|
|
124
|
+
return int(sizes.min())
|
|
125
|
+
|
|
126
|
+
@property
|
|
127
|
+
def max_size(self) -> int:
|
|
128
|
+
"""Maximum neighborhood size."""
|
|
129
|
+
sizes = np.diff(self.adjacency.indptr)
|
|
130
|
+
return int(sizes.max())
|
|
131
|
+
|
|
132
|
+
def __repr__(self) -> str:
|
|
133
|
+
return (
|
|
134
|
+
f"_SphereNeighborhoods(n_voxels={self.n_voxels}, "
|
|
135
|
+
f"radius={self.radius}mm, "
|
|
136
|
+
f"mean_size={self.mean_size:.1f})"
|
|
137
|
+
)
|
|
138
|
+
|
|
139
|
+
|
|
140
|
+
def compute_searchlight_neighborhoods(
|
|
141
|
+
mask_img: Nifti1Image,
|
|
142
|
+
radius: float = 10.0,
|
|
143
|
+
) -> _SphereNeighborhoods:
|
|
144
|
+
"""Compute sphere neighborhoods for all voxels in a brain mask.
|
|
145
|
+
|
|
146
|
+
For each voxel in the mask, this function identifies all other voxels
|
|
147
|
+
within the specified radius (in millimeters).
|
|
148
|
+
|
|
149
|
+
The algorithm uses sklearn's BallTree for efficient radius queries in
|
|
150
|
+
world coordinates (mm), ensuring accurate neighborhoods regardless of
|
|
151
|
+
voxel resolution.
|
|
152
|
+
|
|
153
|
+
Args:
|
|
154
|
+
mask_img: NIfTI mask image defining the brain region
|
|
155
|
+
radius: Radius of spheres in millimeters (default: 10.0)
|
|
156
|
+
|
|
157
|
+
Returns:
|
|
158
|
+
_SphereNeighborhoods with precomputed adjacency matrix
|
|
159
|
+
|
|
160
|
+
Raises:
|
|
161
|
+
ValueError: If mask has no non-zero voxels
|
|
162
|
+
|
|
163
|
+
Examples:
|
|
164
|
+
```python
|
|
165
|
+
import nibabel as nib
|
|
166
|
+
|
|
167
|
+
mask = nib.load("brain_mask.nii.gz")
|
|
168
|
+
neighborhoods = compute_searchlight_neighborhoods(mask, radius=8.0)
|
|
169
|
+
|
|
170
|
+
print(neighborhoods)
|
|
171
|
+
# _SphereNeighborhoods(n_voxels=50000, radius=8.0mm, mean_size=33.2)
|
|
172
|
+
```
|
|
173
|
+
"""
|
|
174
|
+
from nilearn.image.resampling import coord_transform
|
|
175
|
+
|
|
176
|
+
mask_data = mask_img.get_fdata().astype(bool)
|
|
177
|
+
affine = mask_img.affine
|
|
178
|
+
|
|
179
|
+
# Get voxel coordinates in world space (mm)
|
|
180
|
+
mask_coords_voxel = np.array(np.nonzero(mask_data)).T # (n_voxels, 3)
|
|
181
|
+
n_voxels = mask_coords_voxel.shape[0]
|
|
182
|
+
|
|
183
|
+
if n_voxels == 0:
|
|
184
|
+
raise ValueError("Mask contains no non-zero voxels")
|
|
185
|
+
|
|
186
|
+
# Transform to world coordinates using affine
|
|
187
|
+
mask_coords_world = np.array(
|
|
188
|
+
coord_transform(
|
|
189
|
+
mask_coords_voxel[:, 0],
|
|
190
|
+
mask_coords_voxel[:, 1],
|
|
191
|
+
mask_coords_voxel[:, 2],
|
|
192
|
+
affine,
|
|
193
|
+
)
|
|
194
|
+
).T # (n_voxels, 3)
|
|
195
|
+
|
|
196
|
+
# Use BallTree for efficient radius queries
|
|
197
|
+
# This is the same approach used by nilearn's searchlight
|
|
198
|
+
clf = neighbors.NearestNeighbors(radius=radius, algorithm="ball_tree")
|
|
199
|
+
clf.fit(mask_coords_world)
|
|
200
|
+
adjacency = clf.radius_neighbors_graph(mask_coords_world, mode="connectivity")
|
|
201
|
+
adjacency = adjacency.tocsr()
|
|
202
|
+
|
|
203
|
+
return _SphereNeighborhoods(
|
|
204
|
+
adjacency=adjacency,
|
|
205
|
+
radius=radius,
|
|
206
|
+
n_voxels=n_voxels,
|
|
207
|
+
)
|