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,501 @@
1
+ """Glass-brain, slice, flatmap, timeseries, and histogram plots for `BrainData`."""
2
+
3
+ import os
4
+ import warnings
5
+
6
+ import numpy as np
7
+
8
+ from nltools.utils import _find_stack_level
9
+ from .utils import _result_from_array
10
+
11
+
12
+ DEFAULT_SLICE_CUT_COORDS = {
13
+ "x": list(range(-50, 51, 8)),
14
+ "y": list(range(-80, 50, 10)),
15
+ "z": list(range(-40, 71, 9)),
16
+ }
17
+
18
+
19
+ def _image_world_bounds(nifti_img, axis_letter: str) -> tuple[float, float]:
20
+ """World-coord (lo, hi) bounds of ``nifti_img`` along ``axis_letter``.
21
+
22
+ Used to bounds-trim ``DEFAULT_SLICE_CUT_COORDS`` so MNI-shaped defaults
23
+ don't trip nilearn's strict cut_coords validation when the data has
24
+ a non-MNI affine or covers a small native-space FOV.
25
+ """
26
+ from nibabel.affines import apply_affine
27
+
28
+ shape = nifti_img.shape[:3]
29
+ affine = nifti_img.affine
30
+ axis_idx = "xyz".index(axis_letter)
31
+ corners = np.array(
32
+ [
33
+ [i, j, k]
34
+ for i in (0, shape[0] - 1)
35
+ for j in (0, shape[1] - 1)
36
+ for k in (0, shape[2] - 1)
37
+ ]
38
+ )
39
+ world = apply_affine(affine, corners)
40
+ return float(world[:, axis_idx].min()), float(world[:, axis_idx].max())
41
+
42
+
43
+ def _plot_brain(
44
+ bd,
45
+ *,
46
+ method="glass",
47
+ upper=None,
48
+ lower=None,
49
+ threshold=None,
50
+ view="z",
51
+ cut_coords=None,
52
+ cmap=None,
53
+ bg_img=None,
54
+ ax=None,
55
+ figsize=(8, 6),
56
+ title=None,
57
+ colorbar=True,
58
+ save=None,
59
+ stat="mean",
60
+ limit=3,
61
+ **kwargs,
62
+ ):
63
+ """Plot BrainData instance using nilearn visualization or matplotlib.
64
+
65
+ Args:
66
+ bd (BrainData): Data to plot.
67
+ method (str): Visualization type ('glass', 'slices', 'timeseries', 'histogram').
68
+ upper (str | float | None): Upper threshold applied to the data
69
+ (nltools semantics; may be a percentile string like ``"95%"``).
70
+ lower (str | float | None): Lower threshold applied to the data
71
+ (nltools semantics).
72
+ threshold (float | str, optional): Absolute-value transparency cutoff
73
+ forwarded to nilearn. Percentile strings such as ``"95%"`` are
74
+ resolved over finite, nonzero magnitudes. Must be >= 0.
75
+ view (str): For ``method="slices"``, any non-empty combination of
76
+ ``"x"``, ``"y"``, ``"z"`` (e.g. ``"xyz"``, ``"xz"``, ``"y"``).
77
+ Default: ``"z"``.
78
+ cut_coords (list or dict, optional): Cut coordinates for multi-slice
79
+ views. If provided, takes precedence over ``view``-based defaults.
80
+ Either a list of per-axis coordinate sequences whose length
81
+ matches ``view``, or a dict keyed by axis letter (``{"x": [...],
82
+ "z": [...]}``) from which entries for each axis in ``view`` are
83
+ looked up.
84
+ cmap (str, optional): Colormap name. By default, positive-only maps use
85
+ ``"Reds"``, negative-only maps use ``"Blues_r"``, and mixed maps
86
+ use ``"RdBu_r"``.
87
+ bg_img (Nifti1Image or str, optional): Background image for slice views.
88
+ ax (matplotlib.axes.Axes, optional): Matplotlib axis to plot on.
89
+ figsize (tuple, optional): default figure size if no axis (8, 6)
90
+ title (str, optional): Plot title.
91
+ colorbar (bool): Whether to show colorbar. Default: True.
92
+ save (str, optional): Path to save figure(s).
93
+ stat (str): Statistic for timeseries plots. Valid options:
94
+ 'mean', 'median', 'std'.
95
+ limit (int): Maximum number of images to render when ``bd`` contains
96
+ multiple maps and ``method`` is ``"glass"`` or ``"slices"``.
97
+ Default: 3. A warning is emitted if the data has more images than
98
+ ``limit``. Ignored for single-image data and for matplotlib-based
99
+ methods (``"timeseries"``, ``"histogram"``), which already
100
+ aggregate across images.
101
+ **kwargs (dict): Additional arguments forwarded to
102
+ `nilearn.plotting.plot_glass_brain` / `plot_stat_map`.
103
+
104
+ Returns:
105
+ matplotlib.figure.Figure | list[matplotlib.figure.Figure]: For
106
+ single-image data, the figure object (last one created if
107
+ `method="slices"` produced multiple per-axis figures). For
108
+ multi-image data with `method` in `{"glass", "slices"}`, a list of
109
+ figures (one per image for glass; one per image-and-view pair for
110
+ slices). All figures auto-display in notebooks.
111
+ """
112
+ import matplotlib.pyplot as plt
113
+ from nilearn.plotting import plot_glass_brain, plot_stat_map
114
+
115
+ from nltools.templates import _get_bg_image
116
+
117
+ # Validate inputs
118
+ if bd.is_empty:
119
+ raise ValueError("Cannot plot empty BrainData object")
120
+
121
+ if threshold is not None and not isinstance(threshold, str) and threshold < 0:
122
+ raise ValueError(
123
+ f"`threshold` is an absolute-value cutoff and must be >= 0 "
124
+ f"(got {threshold}). Use `upper` / `lower` for one-sided data "
125
+ f"thresholding."
126
+ )
127
+
128
+ # Validate 'method' parameter
129
+ valid_methods = ["glass", "slices", "timeseries", "histogram"]
130
+ if method not in valid_methods:
131
+ raise ValueError(
132
+ f"Invalid 'method' parameter: '{method}'. Must be one of: {valid_methods}. "
133
+ )
134
+
135
+ # Handle matplotlib-based plots (timeseries, histogram)
136
+ if method in ["timeseries", "histogram"]:
137
+ return _plot_matplotlib(
138
+ bd, method=method, stat=stat, ax=ax, figsize=figsize, title=title, save=save
139
+ )
140
+
141
+ # Parse `view` into an ordered list of axis letters (only matters for
142
+ # method="slices"; cheap to compute so we always do it).
143
+ views = list(view.lower()) if isinstance(view, str) else []
144
+ if not views or not set(views).issubset({"x", "y", "z"}):
145
+ raise ValueError(
146
+ f"Invalid `view`: {view!r}. Must be a non-empty string containing "
147
+ "any combination of 'x', 'y', 'z' (e.g. 'xyz', 'xz', 'y')."
148
+ )
149
+
150
+ # Resolve cut_coords against `views`. User-supplied cut_coords take
151
+ # precedence; defaults are drawn from DEFAULT_SLICE_CUT_COORDS per-axis.
152
+ # Track whether we picked the defaults so we can bounds-trim them later
153
+ # against the actual image — MNI-shaped defaults land outside the data
154
+ # for native-space or small synthetic maps, and nilearn 0.12 rejects
155
+ # all-out-of-bounds cut_coords with an opaque ValueError.
156
+ cut_coords_were_defaulted = cut_coords is None
157
+ if cut_coords is None:
158
+ cut_coords = [DEFAULT_SLICE_CUT_COORDS[v] for v in views]
159
+ elif isinstance(cut_coords, dict):
160
+ missing = [v for v in views if v not in cut_coords]
161
+ if missing:
162
+ raise ValueError(
163
+ f"`cut_coords` dict is missing entries for axes {missing} "
164
+ f"required by view={view!r}."
165
+ )
166
+ cut_coords = [cut_coords[v] for v in views]
167
+ else:
168
+ cut_coords = [list(c) if isinstance(c, range) else c for c in cut_coords]
169
+ if len(cut_coords) != len(views):
170
+ raise ValueError(
171
+ f"`cut_coords` has {len(cut_coords)} entries but view={view!r} "
172
+ f"requires {len(views)}."
173
+ )
174
+
175
+ # Decide which images to plot. For multi-image data we render up to
176
+ # `limit` maps and return a list of figures so the user can see (and
177
+ # programmatically access) each map instead of silently dropping all
178
+ # but the first.
179
+ multi = len(bd.shape) > 1 and bd.shape[0] > 1
180
+ if multi:
181
+ n_total = bd.shape[0]
182
+ n_to_plot = min(n_total, limit)
183
+ if n_total > limit:
184
+ warnings.warn(
185
+ f"BrainData contains {n_total} images; plotting first "
186
+ f"{n_to_plot}. Pass `limit={n_total}` (or higher) to plot "
187
+ "more, or index/aggregate before calling .plot().",
188
+ UserWarning,
189
+ stacklevel=_find_stack_level(),
190
+ )
191
+ sub_objs = [bd[i] for i in range(n_to_plot)]
192
+ else:
193
+ sub_objs = [bd]
194
+
195
+ # Standard-space gate. Glass brain draws an MNI-shape outline, and the
196
+ # default slice background is looked up from the MNI template registry
197
+ # — both are misleading on native-space data. Slices with a user-
198
+ # supplied bg_img work for any space and are the documented escape
199
+ # hatch (see Miyawaki / native-space tutorials).
200
+ from nltools.templates import _is_standard_space
201
+
202
+ standard, reason = _is_standard_space(bd.mask.affine)
203
+ if not standard:
204
+ if method == "glass":
205
+ if bg_img is not None:
206
+ warnings.warn(
207
+ f"method='glass' requires standard MNI space ({reason}); "
208
+ f"falling back to method='slices' with the bg_img you "
209
+ "provided.",
210
+ UserWarning,
211
+ stacklevel=_find_stack_level(),
212
+ )
213
+ method = "slices"
214
+ else:
215
+ raise ValueError(
216
+ f"method='glass' requires data in standard MNI space, "
217
+ f"but {reason}. Pass method='slices' with "
218
+ f"bg_img=<your subject anatomical>, or call "
219
+ f"bd.resample() to bring data into standard space first."
220
+ )
221
+ elif method == "slices" and bg_img is None:
222
+ raise ValueError(
223
+ f"Cannot auto-resolve a background image for non-standard-"
224
+ f"space data ({reason}). Pass "
225
+ f"bg_img=<your subject anatomical>, or call bd.resample() "
226
+ f"to bring data into standard space first."
227
+ )
228
+
229
+ # Resolve background image once for slices (template lookup is the same
230
+ # across images sharing a mask).
231
+ if method == "slices" and bg_img is None:
232
+ bg_img = _get_bg_image(bd.mask.affine)
233
+
234
+ # Collect the matplotlib figure underlying each nilearn display, so the
235
+ # return value has a standard `_repr_*_` path and is recognized by
236
+ # frontend filters. For single-image data the last figure is detached
237
+ # from pyplot below to avoid double-display via `flush_figures`.
238
+ figures = []
239
+
240
+ for idx, sub in enumerate(sub_objs):
241
+ # Apply thresholding per-image
242
+ if upper is not None or lower is not None:
243
+ obj = sub.threshold(upper=upper, lower=lower)
244
+ else:
245
+ obj = sub
246
+
247
+ from .utils import _resolve_threshold
248
+
249
+ threshold_use = _resolve_threshold(threshold, np.abs(obj.data))
250
+ if threshold_use is not None and threshold_use < 0:
251
+ raise ValueError(
252
+ f"`threshold` is an absolute-value cutoff and must be >= 0 "
253
+ f"(got {threshold_use}). Use `upper` / `lower` for one-sided "
254
+ f"data thresholding."
255
+ )
256
+ displayed_data = obj.data
257
+ if threshold_use is not None:
258
+ displayed_data = displayed_data[np.abs(displayed_data) >= threshold_use]
259
+ cmap_use = cmap if cmap is not None else _auto_select_colormap(displayed_data)
260
+ save_paths = _prepare_save_paths(save, idx if multi else None) if save else None
261
+
262
+ # A plot cannot show NaN/inf; nilearn zero-fills them itself but warns
263
+ # every time, which is noise for ROI maps and tSNR (NaN outside parcels
264
+ # or where std == 0). Zero-fill up front so the result is identical
265
+ # and silent.
266
+ if not np.all(np.isfinite(obj.data)):
267
+ obj = _result_from_array(
268
+ obj,
269
+ np.nan_to_num(obj.data, nan=0.0, posinf=0.0, neginf=0.0),
270
+ rows="preserve",
271
+ )
272
+
273
+ try:
274
+ nifti_img = obj.to_nifti()
275
+ except Exception as e:
276
+ raise RuntimeError(f"Failed to convert BrainData to NIfTI: {e}") from e
277
+
278
+ # Per-image kwargs. Use the BrainData mask as a transparency image
279
+ # so voxels outside the mask render transparent (nilearn >= 0.12).
280
+ # Users can override by passing their own `transparency=` kwarg.
281
+ plot_kwargs = kwargs.copy()
282
+ plot_kwargs.pop("how", None)
283
+ if multi:
284
+ sub_title = f"{title} (image {idx})" if title else f"image {idx}"
285
+ else:
286
+ sub_title = title
287
+ if sub_title:
288
+ plot_kwargs["title"] = sub_title
289
+ if threshold_use is not None:
290
+ plot_kwargs["threshold"] = threshold_use
291
+ plot_kwargs.setdefault("transparency", obj.mask)
292
+
293
+ if method == "glass":
294
+ display_glass = plot_glass_brain(
295
+ nifti_img,
296
+ display_mode="lzry",
297
+ colorbar=colorbar,
298
+ cmap=cmap_use,
299
+ plot_abs=False,
300
+ **plot_kwargs,
301
+ )
302
+ fig = display_glass.frame_axes.figure
303
+ if save_paths:
304
+ fig.savefig(save_paths["glass"], bbox_inches="tight")
305
+ figures.append(fig)
306
+
307
+ elif method == "slices":
308
+ for v, c in zip(views, cut_coords):
309
+ savefile = save_paths["slices"][v] if save_paths else None
310
+ # Trim defaulted MNI cut_coords down to coords that actually
311
+ # land inside the image. If nothing remains (small native-
312
+ # space FOV, single-slice synthetic data), pass None so
313
+ # nilearn picks its own coords within the bounds.
314
+ c_eff = c
315
+ if cut_coords_were_defaulted and isinstance(c_eff, list):
316
+ lo, hi = _image_world_bounds(nifti_img, v)
317
+ in_bounds = [float(x) for x in c_eff if lo <= float(x) <= hi]
318
+ c_eff = in_bounds if in_bounds else None
319
+ display_slice = plot_stat_map(
320
+ nifti_img,
321
+ cut_coords=c_eff,
322
+ display_mode=v,
323
+ cmap=cmap_use,
324
+ bg_img=bg_img,
325
+ colorbar=colorbar,
326
+ **plot_kwargs,
327
+ )
328
+ fig = display_slice.frame_axes.figure
329
+ if savefile:
330
+ fig.savefig(savefile, bbox_inches="tight")
331
+ figures.append(fig)
332
+
333
+ if not figures:
334
+ return None
335
+ if multi:
336
+ # Leave all figures attached to pyplot's tracker so notebook
337
+ # auto-display via `flush_figures` renders each one. Return the
338
+ # list for programmatic access.
339
+ return figures
340
+ # Single image: detach only the figure we return so its `_repr_*_`
341
+ # rendering doesn't duplicate via `flush_figures`. Any earlier per-view
342
+ # figures from method="slices" stay on pyplot's tracker so the cell's
343
+ # post-hook can display them.
344
+ plt.close(figures[-1])
345
+ return figures[-1]
346
+
347
+
348
+ def _plot_matplotlib(
349
+ bd, method, stat="mean", figsize=(8, 6), ax=None, title=None, save=None
350
+ ):
351
+ """Plot using matplotlib (timeseries or histogram).
352
+
353
+ Args:
354
+ bd (BrainData): Data to plot.
355
+ method (str): 'timeseries' or 'histogram'.
356
+ stat (str): Statistic for timeseries ('mean', 'median', 'std').
357
+ figsize (tuple): Figure size when no axis is given. Default: (8, 6).
358
+ ax (matplotlib.axes.Axes | None): Existing axis to plot on.
359
+ title (str | None): Plot title.
360
+ save (str | None): Path to save the figure.
361
+
362
+ Returns:
363
+ matplotlib.figure.Figure: The rendered figure.
364
+ """
365
+ import matplotlib.pyplot as plt
366
+
367
+ # Create axis if not provided. Track ownership so we only detach figures
368
+ # we created from pyplot's tracker — caller-supplied axes belong to the
369
+ # caller's figure lifecycle.
370
+ if ax is None:
371
+ fig, ax = plt.subplots(figsize=figsize)
372
+ owns_fig = True
373
+ else:
374
+ fig = ax.figure
375
+ owns_fig = False
376
+
377
+ if method == "timeseries":
378
+ # For single image, raise informative error
379
+ if len(bd.shape) == 1 or (len(bd.shape) > 1 and bd.shape[0] == 1):
380
+ raise ValueError(
381
+ "timeseries plotting requires multiple images. "
382
+ f"Got {bd.shape[0] if len(bd.shape) > 1 else 1} image(s). "
383
+ "Use histogram for single image visualization."
384
+ )
385
+
386
+ # Compute statistic across voxels for each image
387
+ if stat == "mean":
388
+ values = bd.mean(axis=1)
389
+ elif stat == "median":
390
+ values = bd.median(axis=1)
391
+ elif stat == "std":
392
+ values = bd.std(axis=1)
393
+ else:
394
+ raise ValueError(
395
+ f"Invalid stat '{stat}'. Must be 'mean', 'median', or 'std'"
396
+ )
397
+
398
+ # Ensure values is 1D array
399
+ if hasattr(values, "data"):
400
+ values = values.data
401
+ values = np.array(values).flatten()
402
+
403
+ # Plot
404
+ ax.plot(values, linewidth=2)
405
+ ax.set_xlabel("Image Index", fontsize=12)
406
+ ax.set_ylabel(f"{stat.capitalize()} Across Voxels", fontsize=12)
407
+ if title is None:
408
+ title = f"{stat.capitalize()} Across Voxels"
409
+ ax.set_title(title, fontsize=14)
410
+ ax.grid(True, alpha=0.3)
411
+
412
+ elif method == "histogram":
413
+ # Flatten data for histogram
414
+ if len(bd.shape) == 1:
415
+ data_flat = bd.data
416
+ else:
417
+ data_flat = bd.data.flatten()
418
+
419
+ # Remove NaN/Inf
420
+ data_flat = data_flat[np.isfinite(data_flat)]
421
+
422
+ # Plot histogram
423
+ ax.hist(data_flat, bins=50, edgecolor="black", alpha=0.7)
424
+ ax.set_xlabel("Voxel Value", fontsize=12)
425
+ ax.set_ylabel("Frequency", fontsize=12)
426
+ if title is None:
427
+ title = "Voxel Value Distribution"
428
+ ax.set_title(title, fontsize=14)
429
+ ax.grid(True, alpha=0.3)
430
+
431
+ # Save if requested
432
+ if save:
433
+ fig.savefig(save, bbox_inches="tight", dpi=150)
434
+
435
+ if owns_fig:
436
+ plt.close(fig)
437
+ return fig
438
+
439
+
440
+ def _auto_select_colormap(data):
441
+ """Auto-select colormap based on data characteristics.
442
+
443
+ Args:
444
+ data (np.ndarray): Brain data values.
445
+
446
+ Returns:
447
+ str: ``'Reds'`` for positive-only data, ``'Blues_r'`` for
448
+ negative-only data, otherwise ``'RdBu_r'``.
449
+ """
450
+ # Flatten data for analysis
451
+ if data.ndim > 1:
452
+ data_flat = data.flatten()
453
+ else:
454
+ data_flat = data
455
+
456
+ # Stored zeros are background, not evidence that a map is mixed-signed.
457
+ data_flat = data_flat[np.isfinite(data_flat) & (data_flat != 0)]
458
+
459
+ if len(data_flat) == 0:
460
+ return "RdBu_r" # Default fallback
461
+
462
+ if np.all(data_flat > 0):
463
+ return "Reds"
464
+ if np.all(data_flat < 0):
465
+ return "Blues_r"
466
+ return "RdBu_r"
467
+
468
+
469
+ def _prepare_save_paths(save, idx=None):
470
+ """Prepare save paths for multiple plot outputs.
471
+
472
+ Args:
473
+ save (str | Path): Base save path; its extension is reused (default
474
+ `png`).
475
+ idx (int | None): Image index appended as ``_img{idx}`` to the base
476
+ filename, to disambiguate saves across multiple images.
477
+
478
+ Returns:
479
+ dict: `'glass'` maps to one path; `'slices'` maps to a dict of per-axis
480
+ (`'x'`, `'y'`, `'z'`) paths.
481
+ """
482
+ save = str(save) # Convert Path objects to strings
483
+ path, filename = os.path.split(save)
484
+ if "." in filename:
485
+ filename, extension = filename.rsplit(".", 1)
486
+ else:
487
+ extension = "png"
488
+
489
+ if idx is not None:
490
+ filename = f"{filename}_img{idx}"
491
+
492
+ base_path = os.path.join(path, filename) if path else filename
493
+
494
+ return {
495
+ "glass": f"{base_path}_glass.{extension}",
496
+ "slices": {
497
+ "x": f"{base_path}_x.{extension}",
498
+ "y": f"{base_path}_y.{extension}",
499
+ "z": f"{base_path}_z.{extension}",
500
+ },
501
+ }