Skip to content

Commit a091c07

Browse files
committed
ENH: Allow selecting projectors for reconstruction
1 parent 0363ac5 commit a091c07

2 files changed

Lines changed: 163 additions & 9 deletions

File tree

mne/_fiff/proj.py

Lines changed: 56 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -515,19 +515,61 @@ def plot_projs_topomap(
515515
)
516516
return fig
517517

518-
def _reconstruct_proj(self, mode="accurate", origin="auto"):
518+
def _reconstruct_proj(self, projs=None, mode="accurate", origin="auto"):
519519
from ..forward import _map_meg_or_eeg_channels
520520

521-
if len(self.info["projs"]) == 0:
522-
return self
523-
self.apply_proj()
521+
if projs is None:
522+
if len(self.info["projs"]) == 0:
523+
return self
524+
self.apply_proj()
525+
info = self.info
526+
else:
527+
if isinstance(projs, Projection):
528+
projs = [projs]
529+
selected_projs = _check_projs(projs)
530+
if len(selected_projs) == 0:
531+
return self
532+
projector, nproj, _ = make_projector(
533+
selected_projs, self.info["ch_names"], self.info["bads"]
534+
)
535+
if nproj == 0:
536+
return self
537+
self.data[:] = np.matmul(projector, self.data)
538+
active_eeg_ref_projs = [
539+
proj
540+
for proj in self.info["projs"]
541+
if proj["active"] and _is_eeg_average_ref_proj(proj)
542+
]
543+
attached_selected = [
544+
attached
545+
for attached in self.info["projs"]
546+
if any(
547+
_proj_equal(attached, selected, check_active=False)
548+
for selected in selected_projs
549+
)
550+
]
551+
activate_proj(attached_selected, copy=False, verbose=False)
552+
mapping_projs = _check_projs([*selected_projs, *active_eeg_ref_projs])
553+
mapping_projs = _uniquify_projs(
554+
mapping_projs, check_active=False, sort=False
555+
)
556+
mapping_projs = activate_proj(mapping_projs, copy=False, verbose=False)
557+
info = self.info.copy()
558+
with info._unlock():
559+
info["projs"] = mapping_projs
524560
for kind in ("meg", "eeg"):
525561
kwargs = dict(meg=False)
526562
kwargs[kind] = True
527-
picks = pick_types(self.info, **kwargs)
563+
picks = pick_types(info, **kwargs)
528564
if len(picks) == 0:
529565
continue
530-
info_from = pick_info(self.info, picks)
566+
info_from = pick_info(info, picks)
567+
if projs is not None:
568+
_, nproj, _ = make_projector(
569+
selected_projs, info_from["ch_names"], info_from["bads"]
570+
)
571+
if nproj == 0:
572+
continue
531573
info_to = info_from.copy()
532574
with info_to._unlock():
533575
info_to["projs"] = []
@@ -1082,9 +1124,7 @@ def _has_eeg_average_ref_proj(
10821124
return False
10831125
found_names = list()
10841126
for proj in projs:
1085-
if proj["kind"] == FIFF.FIFFV_PROJ_ITEM_EEG_AVREF or re.match(
1086-
"^Average .* reference$", proj["desc"]
1087-
):
1127+
if _is_eeg_average_ref_proj(proj):
10881128
if not check_active or proj["active"]:
10891129
found_names.extend(proj["data"]["col_names"])
10901130
# If some are missing we have a problem (keep order for the message,
@@ -1097,6 +1137,13 @@ def _has_eeg_average_ref_proj(
10971137
return True
10981138

10991139

1140+
def _is_eeg_average_ref_proj(proj):
1141+
"""Check whether a projector is an average EEG reference."""
1142+
return proj["kind"] == FIFF.FIFFV_PROJ_ITEM_EEG_AVREF or bool(
1143+
re.match("^Average .* reference$", proj["desc"])
1144+
)
1145+
1146+
11001147
def _needs_eeg_average_ref_proj(info):
11011148
"""Determine if the EEG needs an average EEG reference.
11021149

mne/tests/test_proj.py

Lines changed: 107 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -261,6 +261,113 @@ def test_compute_proj_epochs(tmp_path):
261261
write_proj(fname, ["foo"], overwrite=True)
262262

263263

264+
@pytest.mark.parametrize(
265+
("ch_type", "picks", "proj_kwargs"),
266+
[
267+
("meg", [0, 1, 2, 3, 4, 6, 7], dict(n_grad=2, n_mag=2, n_eeg=0)),
268+
("eeg", np.arange(340, 360), dict(n_grad=0, n_mag=0, n_eeg=2)),
269+
],
270+
)
271+
def test_reconstruct_proj_selection(ch_type, picks, proj_kwargs):
272+
"""Test selecting projectors for reconstruction."""
273+
raw = read_raw_fif(raw_fname)
274+
raw.add_proj([], remove_existing=True)
275+
events = read_events(event_fname)
276+
evoked = Epochs(
277+
raw,
278+
events[:5],
279+
1,
280+
-0.1,
281+
0.1,
282+
picks=picks,
283+
decim=10,
284+
preload=True,
285+
verbose="error",
286+
).average()
287+
if ch_type == "eeg":
288+
evoked.set_eeg_reference(projection=True)
289+
evoked_car_active = evoked.copy().apply_proj()
290+
evoked_proj = evoked.copy().apply_proj().crop(None, 0)
291+
projs = compute_proj_evoked(evoked_proj, **proj_kwargs)
292+
assert len(projs) >= 2
293+
evoked.add_proj(projs)
294+
original_projs = cp.deepcopy(evoked.info["projs"])
295+
296+
# The default is unchanged, including activating all attached projectors.
297+
default = evoked.copy()._reconstruct_proj()
298+
default_none = evoked.copy()._reconstruct_proj(projs=None)
299+
assert_allclose(default_none.data, default.data)
300+
assert all(proj["active"] for proj in default.info["projs"])
301+
302+
# A scalar projector and a list of projectors use exactly the requested
303+
# artifact directions, ignoring unrelated projectors attached to Info.
304+
selected = list(projs[:2])
305+
reconstructed = []
306+
for proj in selected:
307+
actual = evoked.copy()._reconstruct_proj(projs=proj)
308+
expected = evoked.copy().del_proj().add_proj(proj)._reconstruct_proj()
309+
assert_allclose(actual.data, expected.data)
310+
assert [p["active"] for p in actual.info["projs"]] == [
311+
p["desc"] == proj["desc"] for p in original_projs
312+
]
313+
reconstructed.append(actual.data)
314+
assert np.linalg.norm(reconstructed[0] - reconstructed[1]) > 1e-3 * np.linalg.norm(
315+
reconstructed[0]
316+
)
317+
318+
actual = evoked.copy()._reconstruct_proj(projs=selected)
319+
expected = evoked.copy().del_proj().add_proj(selected)._reconstruct_proj()
320+
assert_allclose(actual.data, expected.data)
321+
selected_descs = {proj["desc"] for proj in selected}
322+
assert [p["active"] for p in actual.info["projs"]] == [
323+
p["desc"] in selected_descs for p in original_projs
324+
]
325+
assert np.linalg.norm(actual.data - reconstructed[0]) > 1e-3 * np.linalg.norm(
326+
reconstructed[0]
327+
)
328+
329+
if ch_type == "eeg":
330+
# An inactive average-reference projector is not part of the explicit
331+
# selection. An already-active one remains part of the mapping state.
332+
expected = evoked_car_active.copy().add_proj(projs[0])._reconstruct_proj()
333+
evoked_car_active.add_proj(projs)
334+
actual = evoked_car_active._reconstruct_proj(projs=projs[0])
335+
assert_allclose(actual.data, expected.data)
336+
assert_allclose(actual.data.mean(axis=0), 0.0, atol=1e-20)
337+
assert [p["active"] for p in actual.info["projs"]] == [True, True, False]
338+
339+
340+
def test_reconstruct_proj_selection_mixed():
341+
"""Test that explicit reconstruction leaves unaffected channel types alone."""
342+
raw = read_raw_fif(raw_fname)
343+
raw.add_proj([], remove_existing=True)
344+
events = read_events(event_fname)
345+
picks = [0, 1, 2, 3, 4, 6, 7, *range(340, 360)]
346+
evoked = Epochs(
347+
raw,
348+
events[:5],
349+
1,
350+
-0.1,
351+
0.1,
352+
picks=picks,
353+
decim=10,
354+
preload=True,
355+
verbose="error",
356+
).average()
357+
projs = compute_proj_evoked(
358+
evoked.copy().pick("meg").crop(None, 0), n_grad=2, n_mag=2, n_eeg=0
359+
)
360+
evoked.add_proj(projs)
361+
eeg_picks = pick_types(evoked.info, meg=False, eeg=True)
362+
meg_picks = pick_types(evoked.info, meg=True, eeg=False)
363+
data_before = evoked.data.copy()
364+
evoked._reconstruct_proj(projs=projs[0])
365+
assert_allclose(evoked.data[eeg_picks], data_before[eeg_picks], rtol=0, atol=0)
366+
assert not np.array_equal(evoked.data[meg_picks], data_before[meg_picks])
367+
assert evoked.info["projs"][0]["active"]
368+
assert not any(proj["active"] for proj in evoked.info["projs"][1:])
369+
370+
264371
@pytest.mark.slowtest
265372
def test_compute_proj_raw(tmp_path):
266373
"""Test SSP computation on raw."""

0 commit comments

Comments
 (0)