Skip to content

Commit e3b2f0a

Browse files
committed
fixed linting errors that came with new range of python versions
1 parent 7a59639 commit e3b2f0a

37 files changed

Lines changed: 2978 additions & 185 deletions

invert/adapters/context_lstm.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,8 +1,9 @@
11
from __future__ import annotations
22

33
import logging
4+
from collections.abc import Callable
45
from copy import deepcopy
5-
from typing import TYPE_CHECKING, Any, Callable
6+
from typing import TYPE_CHECKING, Any
67

78
import numpy as np
89

invert/auto_inverse/auto_inverse.py

Lines changed: 5 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,4 @@
11
import logging
2-
from typing import Optional, Union
32

43
import mne
54
import numpy as np
@@ -88,15 +87,15 @@ def generate_summary(self) -> str:
8887

8988
return summary
9089

91-
def get_recommended_solver(self) -> Optional[str]:
90+
def get_recommended_solver(self) -> str | None:
9291
"""Get the recommended solver name."""
9392
return self.recommended_solver
9493

9594

9695
class AutoInverse:
9796
def __init__(
9897
self,
99-
prior: Union[str, PriorEnum] = "patch",
98+
prior: str | PriorEnum = "patch",
10099
snr="auto",
101100
alpha="auto",
102101
n_samples=100,
@@ -148,16 +147,16 @@ def __init__(
148147
self.n_samples = n_samples
149148
self.n_timepoints = n_timepoints
150149
self.verbose = verbose
151-
self.report: Optional[Report] = None
150+
self.report: Report | None = None
152151
self.sim_params = None
153152
self.solvers_to_test = None
154153
self.simulation_config = None
155154

156155
def fit(
157156
self,
158-
data: Union[EvokedArray, Evoked],
157+
data: EvokedArray | Evoked,
159158
forward: Forward,
160-
) -> Optional[Report]:
159+
) -> Report | None:
161160
"""
162161
Fit the auto inverse model to determine the best solver.
163162

invert/benchmark/datasets.py

Lines changed: 2 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,13 +1,12 @@
1-
from typing import Union
21

32
from pydantic import BaseModel
43

54

65
class DatasetConfig(BaseModel):
76
name: str
87
description: str
9-
n_sources: Union[int, tuple[int, int]]
10-
n_orders: Union[int, tuple[int, int]]
8+
n_sources: int | tuple[int, int]
9+
n_orders: int | tuple[int, int]
1110
snr_range: tuple[float, float]
1211
n_timepoints: int
1312
n_samples: int = 50

invert/benchmark/runner.py

Lines changed: 12 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -87,13 +87,18 @@
8787
"SAM": ("invert.solvers.beamformers.sam", "SolverSAM"),
8888
# "EBB": ("invert.solvers.beamformers.ebb", "SolverEBB"),
8989
"ESMV": ("invert.solvers.beamformers.esmv", "SolverESMV"),
90+
"ESMV-MVPURE": ("invert.solvers.beamformers.esmv_mvpure", "SolverESMVMVPURE"),
9091
"ESMV2": ("invert.solvers.beamformers.esmv2", "SolverESMV2"),
9192
"ESMV3": ("invert.solvers.beamformers.esmv3", "SolverESMV3"),
9293
"DeblurFlexESMV": (
9394
"invert.solvers.beamformers.deblur_flex_esmv",
9495
"SolverDeblurFlexESMV",
9596
),
9697
"FlexESMV": ("invert.solvers.beamformers.flex_esmv", "SolverFlexESMV"),
98+
"FlexESMV-MVPURE": (
99+
"invert.solvers.beamformers.flex_esmv_mvpure",
100+
"SolverFlexESMVMVPURE",
101+
),
97102
"SafeFlexESMV": ("invert.solvers.beamformers.safe_flex_esmv", "SolverSafeFlexESMV"),
98103
"SharpFlexESMV": (
99104
"invert.solvers.beamformers.sharp_flex_esmv",
@@ -226,12 +231,16 @@ def _expects_simulation_config(solver_cls: type[BaseSolver]) -> bool:
226231

227232

228233
def _default_nn_batch_size(forward: mne.Forward) -> int:
229-
"""Default ANN training batch size: 2x number of dipoles in the source model."""
234+
"""Default ANN training batch size: ~12x number of dipoles (min 4096).
235+
236+
Larger batches provide more diverse training data per epoch, which
237+
significantly improves ANN solver quality (see ann_testbench round-2/3).
238+
"""
230239
try:
231240
n_dipoles = int(forward["sol"]["data"].shape[1])
232241
except Exception:
233242
return 1
234-
return max(1, 2 * n_dipoles)
243+
return max(4096, 12 * n_dipoles)
235244

236245

237246
def _generate_solver_categories() -> dict[str, list[str]]:
@@ -579,7 +588,7 @@ def run(self) -> list[BenchmarkResult]:
579588
update={"batch_size": _default_nn_batch_size(self.forward)}
580589
)
581590
logger.info(
582-
"ANN training batch_size=%d (default=2*n_dipoles) for %s",
591+
"ANN training batch_size=%d for %s",
583592
int(train_sim_config.batch_size),
584593
solver_name,
585594
)

invert/benchmark/visualize.py

Lines changed: 4 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,5 @@
11
import json
22
from pathlib import Path
3-
from typing import Optional, Union
43

54
import matplotlib.pyplot as plt
65
import numpy as np
@@ -36,9 +35,9 @@ def _solver_to_category(solver_name: str) -> str:
3635

3736

3837
def visualize_results(
39-
results_path_or_data: Union[str, Path, list[BenchmarkResult]],
40-
metrics: Optional[list[str]] = None,
41-
save_path: Optional[Union[str, Path]] = None,
38+
results_path_or_data: str | Path | list[BenchmarkResult],
39+
metrics: list[str] | None = None,
40+
save_path: str | Path | None = None,
4241
) -> list[plt.Figure]:
4342
if isinstance(results_path_or_data, (str, Path)):
4443
path = Path(results_path_or_data)
@@ -120,7 +119,7 @@ def visualize_results(
120119

121120
# Build legend grouped by category, placed below the plot
122121
handles, labels = ax.get_legend_handles_labels()
123-
label_to_handle = dict(zip(labels, handles))
122+
label_to_handle = dict(zip(labels, handles, strict=False))
124123

125124
legend_handles = []
126125
legend_labels = []

invert/config.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -43,17 +43,20 @@
4343
"Source-MAP-MSP",
4444
"MVAB",
4545
"LCMV",
46+
"LCMV-MVPURE",
4647
"DICS",
4748
"SMV",
4849
"WNMV",
4950
"HOCMV",
5051
"ESMV",
52+
"ESMV-MVPURE",
5153
"MCMV",
5254
"HOCMCMV",
5355
"ReciPSIICOS",
5456
"SAM",
5557
"Adapt-Flex-ESMV",
5658
"Flex-ESMV",
59+
"Flex-ESMV-MVPURE",
5760
"Flex-ESMV2",
5861
"Deblur-Flex-ESMV",
5962
"Safe-Flex-ESMV",

invert/denoising/ASR.py

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,4 @@
11
import logging
2-
from typing import Union
32

43
import mne
54
import numpy as np
@@ -116,7 +115,7 @@ def run(self, data: np.ndarray, calibration=None) -> np.ndarray:
116115

117116
# -------------------------- mne wrapper --------------------------
118117

119-
def run_mne(self, mne_obj: Union[mne.io.BaseRaw, mne.Epochs, mne.Evoked]):
118+
def run_mne(self, mne_obj: mne.io.BaseRaw | mne.Epochs | mne.Evoked):
120119
"""
121120
Run ASR on an MNE object and return a cleaned copy.
122121
"""

invert/denoising/lgsp.py

Lines changed: 10 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -1,13 +1,13 @@
11
from dataclasses import dataclass
2-
from typing import Literal, Optional, Union
2+
from typing import Literal
33

44
import mne
55
import numpy as np
66
from scipy.linalg import eigh
77
from sklearn.covariance import OAS
88

99
ArrayLike = np.ndarray
10-
MNEObj = Union[mne.io.BaseRaw, mne.Epochs, mne.Evoked, mne.EvokedArray]
10+
MNEObj = mne.io.BaseRaw | mne.Epochs | mne.Evoked | mne.EvokedArray
1111

1212

1313
# ============================
@@ -51,12 +51,12 @@ def _inv_sqrt_spd(Q: np.ndarray, d: np.ndarray) -> np.ndarray:
5151

5252
@dataclass
5353
class DenoiserBase:
54-
sfreq: Optional[float] = None
54+
sfreq: float | None = None
5555
window_size: float = 0.5 # seconds
5656
step_frac: float = 0.5 # overlap = 1 - step_frac
5757
use_windows: bool = True # set False to process full signal
5858

59-
def _ensure_sfreq(self, sfreq: Optional[float]):
59+
def _ensure_sfreq(self, sfreq: float | None):
6060
if self.sfreq is None:
6161
if sfreq is None:
6262
raise ValueError(
@@ -65,7 +65,7 @@ def _ensure_sfreq(self, sfreq: Optional[float]):
6565
self.sfreq = sfreq
6666

6767
# ---- public API on NumPy ----
68-
def run(self, X: ArrayLike, sfreq: Optional[float] = None) -> ArrayLike:
68+
def run(self, X: ArrayLike, sfreq: float | None = None) -> ArrayLike:
6969
"""
7070
X shape must be (n_channels, n_times).
7171
"""
@@ -137,12 +137,12 @@ def _process_window(self, Xw: ArrayLike) -> ArrayLike:
137137

138138
@dataclass
139139
class LGSP(DenoiserBase):
140-
L: Optional[np.ndarray] = None # (m, n_sources)
140+
L: np.ndarray | None = None # (m, n_sources)
141141
sigma_prior: Literal["identity"] = "identity"
142142
lambda_ref: float = 1e-3 # ridge on C_ref
143143
alpha_model: float = 0.2 # blend with scaled I: Cref'=(1-a)Cref+a*trace(Cref)/m*I
144-
rank: Optional[int] = None # keep r smallest generalized eigenvalues
145-
tau: Optional[float] = None # or threshold on generalized eigenvalues
144+
rank: int | None = None # keep r smallest generalized eigenvalues
145+
tau: float | None = None # or threshold on generalized eigenvalues
146146
center: bool = True
147147
shrink_data: bool = True
148148

@@ -164,7 +164,7 @@ def _build_Cref(self, n_ch: int) -> np.ndarray:
164164
return Cref
165165

166166
def _gevd_brain_projector(
167-
self, Cx: np.ndarray, Cref: np.ndarray, r: Optional[int], tau: Optional[float]
167+
self, Cx: np.ndarray, Cref: np.ndarray, r: int | None, tau: float | None
168168
) -> np.ndarray:
169169
"""
170170
Return the sensor-space projector P onto the 'brain-like' subspace
@@ -214,7 +214,7 @@ def _process_window(self, Xw: ArrayLike) -> ArrayLike:
214214

215215
@dataclass
216216
class SRB(DenoiserBase):
217-
L: Optional[np.ndarray] = None # (m, n_sources) required for meaningful use
217+
L: np.ndarray | None = None # (m, n_sources) required for meaningful use
218218
mu: float = 1e-1 # Tikhonov on sources (||S||^2)
219219
beta: float = 0.5 # blend factor: X_clean = beta*L*W*X + (1-beta)*X
220220
center: bool = True

invert/ensemble/ensemble.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -46,7 +46,7 @@ def apply_inverse_operator(self, mne_obj, *args, **kwargs):
4646

4747
logger.info("Neg Log Likelihoods:")
4848
for solver_name, neg_log_likelihood in zip(
49-
self.solver_names, self.neg_log_likelihoods
49+
self.solver_names, self.neg_log_likelihoods, strict=False
5050
):
5151
logger.info(f"{solver_name}: {neg_log_likelihood}")
5252
logger.info(f"Final likelihood: {final_neg_log_likelihood}")

invert/evaluate/evaluate.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -259,7 +259,7 @@ def nmse(y_true, y_pred):
259259
if np.any(np.isnan(y_true)) or np.any(np.isnan(y_pred)):
260260
return np.nan
261261
error = np.zeros(y_true.shape[1])
262-
for i, (y_true_slice, y_pred_slice) in enumerate(zip(y_true.T, y_pred.T)):
262+
for i, (y_true_slice, y_pred_slice) in enumerate(zip(y_true.T, y_pred.T, strict=False)):
263263
error[i] = np.mean(
264264
(
265265
(y_true_slice / abs(y_true_slice).max())
@@ -274,7 +274,7 @@ def corr(y_true, y_pred):
274274
if np.any(np.isnan(y_true)) or np.any(np.isnan(y_pred)):
275275
return np.nan
276276
error = np.zeros(y_true.shape[1])
277-
for i, (y_true_slice, y_pred_slice) in enumerate(zip(y_true.T, y_pred.T)):
277+
for i, (y_true_slice, y_pred_slice) in enumerate(zip(y_true.T, y_pred.T, strict=False)):
278278
error[i] = pearsonr(y_true_slice, y_pred_slice)[0]
279279
return error
280280

0 commit comments

Comments
 (0)