Skip to content

Commit 3d725a4

Browse files
authored
Merge pull request #3249 from chrishalcrow/simplify-qm-tests
Refactor quality metrics tests to use fixture
2 parents 73f4d58 + 41f73ed commit 3d725a4

3 files changed

Lines changed: 38 additions & 77 deletions

File tree

src/spikeinterface/qualitymetrics/tests/conftest.py

Lines changed: 37 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -5,8 +5,11 @@
55
create_sorting_analyzer,
66
)
77

8+
job_kwargs = dict(n_jobs=2, progress_bar=True, chunk_duration="1s")
89

9-
def _small_sorting_analyzer():
10+
11+
@pytest.fixture(scope="module")
12+
def small_sorting_analyzer():
1013
recording, sorting = generate_ground_truth_recording(
1114
durations=[2.0],
1215
num_units=10,
@@ -33,5 +36,36 @@ def _small_sorting_analyzer():
3336

3437

3538
@pytest.fixture(scope="module")
36-
def small_sorting_analyzer():
37-
return _small_sorting_analyzer()
39+
def sorting_analyzer_simple():
40+
# we need high firing rate for amplitude_cutoff
41+
recording, sorting = generate_ground_truth_recording(
42+
durations=[
43+
120.0,
44+
],
45+
sampling_frequency=30_000.0,
46+
num_channels=6,
47+
num_units=10,
48+
generate_sorting_kwargs=dict(firing_rates=10.0, refractory_period_ms=4.0),
49+
generate_unit_locations_kwargs=dict(
50+
margin_um=5.0,
51+
minimum_z=5.0,
52+
maximum_z=20.0,
53+
),
54+
generate_templates_kwargs=dict(
55+
unit_params=dict(
56+
alpha=(200.0, 500.0),
57+
)
58+
),
59+
noise_kwargs=dict(noise_levels=5.0, strategy="tile_pregenerated"),
60+
seed=1205,
61+
)
62+
63+
sorting_analyzer = create_sorting_analyzer(sorting, recording, format="memory", sparse=True)
64+
65+
sorting_analyzer.compute("random_spikes", max_spikes_per_unit=300, seed=1205)
66+
sorting_analyzer.compute("noise_levels")
67+
sorting_analyzer.compute("waveforms", **job_kwargs)
68+
sorting_analyzer.compute("templates")
69+
sorting_analyzer.compute("spike_amplitudes", **job_kwargs)
70+
71+
return sorting_analyzer

src/spikeinterface/qualitymetrics/tests/test_metrics_functions.py

Lines changed: 1 addition & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -135,36 +135,6 @@ def test_unit_id_order_independence(small_sorting_analyzer):
135135
assert quality_metrics_2[metric][1] == metric_1_data["#4"]
136136

137137

138-
def _sorting_analyzer_simple():
139-
recording, sorting = generate_ground_truth_recording(
140-
durations=[
141-
50.0,
142-
],
143-
sampling_frequency=30_000.0,
144-
num_channels=6,
145-
num_units=10,
146-
generate_sorting_kwargs=dict(firing_rates=6.0, refractory_period_ms=4.0),
147-
noise_kwargs=dict(noise_levels=5.0, strategy="tile_pregenerated"),
148-
seed=2205,
149-
)
150-
151-
sorting_analyzer = create_sorting_analyzer(sorting, recording, format="memory", sparse=True)
152-
153-
sorting_analyzer.compute("random_spikes", max_spikes_per_unit=300, seed=2205)
154-
sorting_analyzer.compute("noise_levels")
155-
sorting_analyzer.compute("waveforms", **job_kwargs)
156-
sorting_analyzer.compute("templates")
157-
sorting_analyzer.compute("principal_components", n_components=5, mode="by_channel_local", **job_kwargs)
158-
sorting_analyzer.compute("spike_amplitudes", **job_kwargs)
159-
160-
return sorting_analyzer
161-
162-
163-
@pytest.fixture(scope="module")
164-
def sorting_analyzer_simple():
165-
return _sorting_analyzer_simple()
166-
167-
168138
def _sorting_violation():
169139
max_time = 100.0
170140
sampling_frequency = 30000
@@ -576,6 +546,7 @@ def test_calculate_sd_ratio(sorting_analyzer_simple):
576546
test_unit_structure_in_output(_small_sorting_analyzer())
577547

578548
# test_calculate_firing_rate_num_spikes(sorting_analyzer)
549+
579550
# test_calculate_snrs(sorting_analyzer)
580551
# test_calculate_amplitude_cutoff(sorting_analyzer)
581552
# test_calculate_presence_ratio(sorting_analyzer)

src/spikeinterface/qualitymetrics/tests/test_quality_metric_calculator.py

Lines changed: 0 additions & 44 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,6 @@
22
from pathlib import Path
33
import numpy as np
44

5-
65
from spikeinterface.core import (
76
generate_ground_truth_recording,
87
create_sorting_analyzer,
@@ -15,54 +14,11 @@
1514
compute_quality_metrics,
1615
)
1716

18-
1917
job_kwargs = dict(n_jobs=2, progress_bar=True, chunk_duration="1s")
2018

2119

22-
def get_sorting_analyzer(seed=2205):
23-
# we need high firing rate for amplitude_cutoff
24-
recording, sorting = generate_ground_truth_recording(
25-
durations=[
26-
120.0,
27-
],
28-
sampling_frequency=30_000.0,
29-
num_channels=6,
30-
num_units=10,
31-
generate_sorting_kwargs=dict(firing_rates=10.0, refractory_period_ms=4.0),
32-
generate_unit_locations_kwargs=dict(
33-
margin_um=5.0,
34-
minimum_z=5.0,
35-
maximum_z=20.0,
36-
),
37-
generate_templates_kwargs=dict(
38-
unit_params=dict(
39-
alpha=(200.0, 500.0),
40-
)
41-
),
42-
noise_kwargs=dict(noise_levels=5.0, strategy="tile_pregenerated"),
43-
seed=seed,
44-
)
45-
46-
sorting_analyzer = create_sorting_analyzer(sorting, recording, format="memory", sparse=True)
47-
48-
sorting_analyzer.compute("random_spikes", max_spikes_per_unit=300, seed=seed)
49-
sorting_analyzer.compute("noise_levels")
50-
sorting_analyzer.compute("waveforms", **job_kwargs)
51-
sorting_analyzer.compute("templates")
52-
sorting_analyzer.compute("spike_amplitudes", **job_kwargs)
53-
54-
return sorting_analyzer
55-
56-
57-
@pytest.fixture(scope="module")
58-
def sorting_analyzer_simple():
59-
sorting_analyzer = get_sorting_analyzer(seed=2205)
60-
return sorting_analyzer
61-
62-
6320
def test_compute_quality_metrics(sorting_analyzer_simple):
6421
sorting_analyzer = sorting_analyzer_simple
65-
print(sorting_analyzer)
6622

6723
# without PCs
6824
metrics = compute_quality_metrics(

0 commit comments

Comments
 (0)