99from tqdm .auto import tqdm
1010
1111from spikeinterface .core .base import unit_period_dtype
12+ from spikeinterface .core .core_tools import slice_rows
1213from spikeinterface .core .job_tools import fix_job_kwargs
1314from spikeinterface .core .sorting_tools import cast_periods_to_unit_period_dtype , remap_unit_indices_in_vector
1415from spikeinterface .core .sortinganalyzer import register_result_extension , AnalyzerExtension
@@ -539,7 +540,7 @@ def _get_data(self, outputs: str = "by_unit"):
539540 (start_sample_index, end_sample_index) tuples.
540541 """
541542 if outputs == "numpy" :
542- good_periods = self .data ["valid_unit_periods" ].copy ()
543+ good_periods = np . asarray ( self .data ["valid_unit_periods" ]) .copy ()
543544 else :
544545 # by_unit
545546 unit_ids = self .sorting_analyzer .unit_ids
@@ -551,7 +552,7 @@ def _get_data(self, outputs: str = "by_unit"):
551552 for unit_index , unit_id in enumerate (unit_ids ):
552553 periods_dict [unit_id ] = []
553554 unit_mask = good_periods_array ["unit_index" ] == unit_index
554- good_periods_unit_segment = good_periods_array [ segment_mask & unit_mask ]
555+ good_periods_unit_segment = slice_rows ( good_periods_array , segment_mask & unit_mask )
555556 for start , end in good_periods_unit_segment [["start_sample_index" , "end_sample_index" ]]:
556557 periods_dict [unit_id ].append ((start , end ))
557558 good_periods .append (periods_dict )
0 commit comments