Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 26 additions & 0 deletions src/spikeinterface/core/channelsaggregationrecording.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,9 @@ class ChannelsAggregationRecording(BaseRecording):
"""
Class that handles aggregating channels from different recordings, e.g. from different channel groups.

Annotations shared by all the recordings, meaning present in every recording and with the same value
everywhere, are propagated to the aggregated recording. All other annotations are dropped.

Do not use this class directly but use `si.aggregate_channels(...)`

"""
Expand Down Expand Up @@ -97,6 +100,23 @@ def __init__(self, recording_list_or_dict=None, renamed_channel_ids=None, record
for prop_name, prop_values in property_dict.items():
self.set_property(key=prop_name, values=prop_values)

# Propagate the annotations that are shared by the recordings. An annotation is shared when every
# recording carries it and all of them agree on its value. Anything else is dropped, which is the
# same rule used by `UnitsAggregationSorting` in unitsaggregationsorting.py.
for annotation_name in recording_list[0].get_annotation_keys():
if not all(annotation_name in rec.get_annotation_keys() for rec in recording_list):
continue
values = [rec.get_annotation(annotation_name, copy=False) for rec in recording_list]
try:
# `np.array_equal` gives a single bool for scalars, strings and arrays alike. Values it
# cannot compare (e.g. ragged object arrays) raise, and are then treated as not shared.
all_values_are_equal = all(np.array_equal(value, values[0]) for value in values[1:])
except Exception:
all_values_are_equal = False
if all_values_are_equal:
# take a copy so the aggregate does not share mutable state with its first child
self.set_annotation(annotation_name, recording_list[0].get_annotation(annotation_name), overwrite=True)

# Aggregate probe information
all_probegroups = [rec.get_probegroup() for rec in recording_list if rec.has_probe()]
if len(all_probegroups) == len(recording_list):
Expand Down Expand Up @@ -254,6 +274,12 @@ def aggregate_channels(
-------
aggregate_recording: ChannelsAggregationRecording
The aggregated recording object

Notes
-----
Annotations are propagated only when they are shared by all the recordings, meaning present in every
recording and with the same value everywhere. Annotations missing from one recording, or with differing
values, are dropped.
"""

return ChannelsAggregationRecording(recording_list_or_dict, renamed_channel_ids, recording_list)
106 changes: 106 additions & 0 deletions src/spikeinterface/core/tests/test_channelsaggregationrecording.py
Original file line number Diff line number Diff line change
Expand Up @@ -332,5 +332,111 @@ def test_aggregate_channels_split_by_round_trip():
assert recovered_names == {"probe_A", "probe_B"}


def test_aggregate_channels_propagates_shared_annotations():
"""Annotations present in every recording with the same value are propagated (issue #3983)."""
recording1 = generate_recording(num_channels=3, durations=[1.0], set_probe=False)
recording2 = generate_recording(num_channels=2, durations=[1.0], set_probe=False)

recording1.annotate(experimenter="alice", session_id=7)
recording2.annotate(experimenter="alice", session_id=7)

aggregated_recording = aggregate_channels([recording1, recording2])

assert aggregated_recording.get_annotation("experimenter") == "alice"
assert aggregated_recording.get_annotation("session_id") == 7

# `is_filtered` is set by all recordings and agrees, so it must be propagated rather than
# falling back to the `BaseRecording.__init__` default of False
assert recording1.get_annotation("is_filtered") == recording2.get_annotation("is_filtered")
assert aggregated_recording.get_annotation("is_filtered") == recording1.get_annotation("is_filtered")


def test_aggregate_channels_drops_conflicting_annotations():
"""Annotations present everywhere but with different values are not propagated (issue #3983)."""
recording1 = generate_recording(num_channels=3, durations=[1.0], set_probe=False)
recording2 = generate_recording(num_channels=2, durations=[1.0], set_probe=False)

recording1.annotate(experimenter="alice")
recording2.annotate(experimenter="bob")

aggregated_recording = aggregate_channels([recording1, recording2])

assert "experimenter" not in aggregated_recording.get_annotation_keys()


def test_aggregate_channels_drops_partial_annotations():
"""Annotations present in only some recordings are not propagated (issue #3983)."""
recording1 = generate_recording(num_channels=3, durations=[1.0], set_probe=False)
recording2 = generate_recording(num_channels=2, durations=[1.0], set_probe=False)

recording1.annotate(experimenter="alice")

aggregated_recording = aggregate_channels([recording1, recording2])

assert "experimenter" not in aggregated_recording.get_annotation_keys()

# also check the other direction, the loop is seeded from the first recording's keys
aggregated_recording_reversed = aggregate_channels([recording2, recording1])
assert "experimenter" not in aggregated_recording_reversed.get_annotation_keys()


def test_aggregate_channels_with_array_valued_annotations():
"""Array-valued and ragged annotations must never raise during aggregation (issue #3983)."""
recording1 = generate_recording(num_channels=3, durations=[1.0], set_probe=False)
recording2 = generate_recording(num_channels=2, durations=[1.0], set_probe=False)

# equal arrays: shared, so propagated
recording1.annotate(reference_waveform=np.arange(5))
recording2.annotate(reference_waveform=np.arange(5))

# arrays of different shape: not shared, so dropped, and must not raise
recording1.annotate(coverage=np.zeros((2, 3)))
recording2.annotate(coverage=np.zeros((4, 7)))

# ragged object values: not provably equal, so dropped, and must not raise
recording1.annotate(ragged=[np.arange(2), np.arange(3)])
recording2.annotate(ragged=[np.arange(2), np.arange(3)])

aggregated_recording = aggregate_channels([recording1, recording2])

assert np.array_equal(aggregated_recording.get_annotation("reference_waveform"), np.arange(5))
assert "coverage" not in aggregated_recording.get_annotation_keys()
# the ragged annotation may be propagated or dropped, the contract is only that we did not crash
assert aggregated_recording.get_num_channels() == 5


def test_aggregate_channels_annotations_do_not_leak_to_inputs():
"""The aggregated recording must not share mutable annotation values with its children."""
recording1 = generate_recording(num_channels=3, durations=[1.0], set_probe=False)
recording2 = generate_recording(num_channels=2, durations=[1.0], set_probe=False)

recording1.annotate(shared_list=[1, 2, 3])
recording2.annotate(shared_list=[1, 2, 3])

aggregated_recording = aggregate_channels([recording1, recording2])
aggregated_recording.get_annotation("shared_list", copy=False).append(4)

assert recording1.get_annotation("shared_list") == [1, 2, 3]
assert recording2.get_annotation("shared_list") == [1, 2, 3]


def test_aggregate_channels_annotations_preserve_probe_information():
"""Propagating annotations must not disturb the probe metadata of the aggregated recording."""
rec_A = _make_rec_with_named_probe("probe_A", "vendor_X", 0.0)
rec_B = _make_rec_with_named_probe("probe_B", "vendor_Y", 1000.0)

contours = [np.asarray(rec.get_probe().probe_planar_contour) for rec in (rec_A, rec_B)]
assert all(contour is not None and contour.size > 0 for contour in contours)

combined = aggregate_channels([rec_A, rec_B])

probes = combined.get_probes()
assert len(probes) == 2
for probe, expected_contour in zip(probes, contours):
assert np.array_equal(np.asarray(probe.probe_planar_contour), expected_contour)
assert np.array_equal(combined.get_channel_locations()[:8], rec_A.get_channel_locations())
assert np.array_equal(combined.get_channel_locations()[8:], rec_B.get_channel_locations())


if __name__ == "__main__":
test_channelsaggregationrecording()
Loading