Skip to content

Commit 82d62ca

Browse files
authored
Merge pull request #3497 from zm711/merge-qc
Fix dtype of quality metrics before and after merging
2 parents 0b1bf67 + 33feca3 commit 82d62ca

3 files changed

Lines changed: 93 additions & 4 deletions

File tree

src/spikeinterface/qualitymetrics/quality_metric_calculator.py

Lines changed: 19 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,7 @@
1717
_misc_metric_name_to_func,
1818
_possible_pc_metric_names,
1919
qm_compute_name_to_column_names,
20+
column_name_to_column_dtype,
2021
)
2122
from .misc_metrics import _default_params as misc_metrics_params
2223
from .pca_metrics import _default_params as pca_metrics_params
@@ -140,13 +141,20 @@ def _merge_extension_data(
140141
all_unit_ids = new_sorting_analyzer.unit_ids
141142
not_new_ids = all_unit_ids[~np.isin(all_unit_ids, new_unit_ids)]
142143

144+
# this creates a new metrics dictionary, but the dtype for everything will be
145+
# object. So we will need to fix this later after computing metrics
143146
metrics = pd.DataFrame(index=all_unit_ids, columns=old_metrics.columns)
144-
145147
metrics.loc[not_new_ids, :] = old_metrics.loc[not_new_ids, :]
146148
metrics.loc[new_unit_ids, :] = self._compute_metrics(
147149
new_sorting_analyzer, new_unit_ids, verbose, metric_names, **job_kwargs
148150
)
149151

152+
# we need to fix the dtypes after we compute everything because we have nans
153+
# we can iterate through the columns and convert them back to the dtype
154+
# of the original quality dataframe.
155+
for column in old_metrics.columns:
156+
metrics[column] = metrics[column].astype(old_metrics[column].dtype)
157+
150158
new_data = dict(metrics=metrics)
151159
return new_data
152160

@@ -229,10 +237,20 @@ def _compute_metrics(self, sorting_analyzer, unit_ids=None, verbose=False, metri
229237
# add NaN for empty units
230238
if len(empty_unit_ids) > 0:
231239
metrics.loc[empty_unit_ids] = np.nan
240+
# num_spikes is an int and should be 0
241+
if "num_spikes" in metrics.columns:
242+
metrics.loc[empty_unit_ids, ["num_spikes"]] = 0
232243

233244
# we use the convert_dtypes to convert the columns to the most appropriate dtype and avoid object columns
234245
# (in case of NaN values)
235246
metrics = metrics.convert_dtypes()
247+
248+
# we do this because the convert_dtypes infers the wrong types sometimes.
249+
# the actual types for columns can be found in column_name_to_column_dtype dictionary.
250+
for column in metrics.columns:
251+
if column in column_name_to_column_dtype:
252+
metrics[column] = metrics[column].astype(column_name_to_column_dtype[column])
253+
236254
return metrics
237255

238256
def _run(self, verbose=False, **job_kwargs):

src/spikeinterface/qualitymetrics/quality_metric_list.py

Lines changed: 40 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -66,7 +66,11 @@
6666
"amplitude_cutoff": ["amplitude_cutoff"],
6767
"amplitude_median": ["amplitude_median"],
6868
"amplitude_cv": ["amplitude_cv_median", "amplitude_cv_range"],
69-
"synchrony": ["sync_spike_2", "sync_spike_4", "sync_spike_8"],
69+
"synchrony": [
70+
"sync_spike_2",
71+
"sync_spike_4",
72+
"sync_spike_8",
73+
],
7074
"firing_range": ["firing_range"],
7175
"drift": ["drift_ptp", "drift_std", "drift_mad"],
7276
"sd_ratio": ["sd_ratio"],
@@ -79,3 +83,38 @@
7983
"silhouette": ["silhouette"],
8084
"silhouette_full": ["silhouette_full"],
8185
}
86+
87+
# this dict allows us to ensure the appropriate dtype of metrics rather than allow Pandas to infer them
88+
column_name_to_column_dtype = {
89+
"num_spikes": int,
90+
"firing_rate": float,
91+
"presence_ratio": float,
92+
"snr": float,
93+
"isi_violations_ratio": float,
94+
"isi_violations_count": float,
95+
"rp_violations": float,
96+
"rp_contamination": float,
97+
"sliding_rp_violation": float,
98+
"amplitude_cutoff": float,
99+
"amplitude_median": float,
100+
"amplitude_cv_median": float,
101+
"amplitude_cv_range": float,
102+
"sync_spike_2": float,
103+
"sync_spike_4": float,
104+
"sync_spike_8": float,
105+
"firing_range": float,
106+
"drift_ptp": float,
107+
"drift_std": float,
108+
"drift_mad": float,
109+
"sd_ratio": float,
110+
"isolation_distance": float,
111+
"l_ratio": float,
112+
"d_prime": float,
113+
"nn_hit_rate": float,
114+
"nn_miss_rate": float,
115+
"nn_isolation": float,
116+
"nn_unit_id": float,
117+
"nn_noise_overlap": float,
118+
"silhouette": float,
119+
"silhouette_full": float,
120+
}

src/spikeinterface/qualitymetrics/tests/test_quality_metric_calculator.py

Lines changed: 34 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -48,6 +48,33 @@ def test_compute_quality_metrics(sorting_analyzer_simple):
4848
assert "isolation_distance" in metrics.columns
4949

5050

51+
def test_merging_quality_metrics(sorting_analyzer_simple):
52+
53+
sorting_analyzer = sorting_analyzer_simple
54+
55+
metrics = compute_quality_metrics(
56+
sorting_analyzer,
57+
metric_names=None,
58+
qm_params=dict(isi_violation=dict(isi_threshold_ms=2)),
59+
skip_pc_metrics=False,
60+
seed=2205,
61+
)
62+
63+
# sorting_analyzer_simple has ten units
64+
new_sorting_analyzer = sorting_analyzer.merge_units([[0, 1]])
65+
66+
new_metrics = new_sorting_analyzer.get_extension("quality_metrics").get_data()
67+
68+
# we should copy over the metrics after merge
69+
for column in metrics.columns:
70+
assert column in new_metrics.columns
71+
# should copy dtype too
72+
assert metrics[column].dtype == new_metrics[column].dtype
73+
74+
# 10 units vs 9 units
75+
assert len(metrics.index) > len(new_metrics.index)
76+
77+
5178
def test_compute_quality_metrics_recordingless(sorting_analyzer_simple):
5279

5380
sorting_analyzer = sorting_analyzer_simple
@@ -106,10 +133,15 @@ def test_empty_units(sorting_analyzer_simple):
106133
seed=2205,
107134
)
108135

109-
for empty_unit_id in sorting_empty.get_empty_unit_ids():
136+
# num_spikes are ints not nans so we confirm empty units are nans for everything except
137+
# num_spikes which should be 0
138+
nan_containing_columns = [column for column in metrics_empty.columns if column != "num_spikes"]
139+
for empty_unit_ids in sorting_empty.get_empty_unit_ids():
110140
from pandas import isnull
111141

112-
assert np.all(isnull(metrics_empty.loc[empty_unit_id].values))
142+
assert np.all(isnull(metrics_empty.loc[empty_unit_ids, nan_containing_columns].values))
143+
if "num_spikes" in metrics_empty.columns:
144+
assert sum(metrics_empty.loc[empty_unit_ids, ["num_spikes"]]) == 0
113145

114146

115147
# TODO @alessio all theses old test should be moved in test_metric_functions.py or test_pca_metrics()

0 commit comments

Comments
 (0)