diff --git a/src/spikeinterface/core/channelsaggregationrecording.py b/src/spikeinterface/core/channelsaggregationrecording.py index 2589e70ea6..bd60602754 100644 --- a/src/spikeinterface/core/channelsaggregationrecording.py +++ b/src/spikeinterface/core/channelsaggregationrecording.py @@ -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(...)` """ @@ -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): @@ -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) diff --git a/src/spikeinterface/core/tests/test_channelsaggregationrecording.py b/src/spikeinterface/core/tests/test_channelsaggregationrecording.py index 792605fde6..4702acca93 100644 --- a/src/spikeinterface/core/tests/test_channelsaggregationrecording.py +++ b/src/spikeinterface/core/tests/test_channelsaggregationrecording.py @@ -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()