Skip to content

Commit d38dbf4

Browse files
authored
Merge pull request #3588 from h-mayorquin/use_strings_as_ids_in_generators
Use strings as ids in generators
2 parents 6fde997 + 7dea3b2 commit d38dbf4

16 files changed

Lines changed: 89 additions & 29 deletions

File tree

src/spikeinterface/core/basesorting.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -135,7 +135,7 @@ def get_total_duration(self) -> float:
135135

136136
def get_unit_spike_train(
137137
self,
138-
unit_id,
138+
unit_id: str | int,
139139
segment_index: Union[int, None] = None,
140140
start_frame: Union[int, None] = None,
141141
end_frame: Union[int, None] = None,

src/spikeinterface/core/generate.py

Lines changed: 8 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,7 @@
22
import math
33
import warnings
44
import numpy as np
5-
from typing import Literal
5+
from typing import Literal, Optional
66
from math import ceil
77

88
from .basesorting import SpikeVectorSortingSegment
@@ -134,7 +134,7 @@ def generate_sorting(
134134
seed = _ensure_seed(seed)
135135
rng = np.random.default_rng(seed)
136136
num_segments = len(durations)
137-
unit_ids = np.arange(num_units)
137+
unit_ids = [str(idx) for idx in np.arange(num_units)]
138138

139139
spikes = []
140140
for segment_index in range(num_segments):
@@ -1111,7 +1111,7 @@ def __init__(
11111111
11121112
"""
11131113

1114-
unit_ids = np.arange(num_units)
1114+
unit_ids = [str(idx) for idx in np.arange(num_units)]
11151115
super().__init__(sampling_frequency, unit_ids)
11161116

11171117
self.num_units = num_units
@@ -1138,6 +1138,7 @@ def __init__(
11381138
firing_rates=firing_rates,
11391139
refractory_period_seconds=self.refractory_period_seconds,
11401140
seed=segment_seed,
1141+
unit_ids=unit_ids,
11411142
t_start=None,
11421143
)
11431144
self.add_sorting_segment(segment)
@@ -1161,6 +1162,7 @@ def __init__(
11611162
firing_rates: float | np.ndarray,
11621163
refractory_period_seconds: float | np.ndarray,
11631164
seed: int,
1165+
unit_ids: list[str],
11641166
t_start: Optional[float] = None,
11651167
):
11661168
self.num_units = num_units
@@ -1177,7 +1179,8 @@ def __init__(
11771179
self.refractory_period_seconds = np.full(num_units, self.refractory_period_seconds, dtype="float64")
11781180

11791181
self.segment_seed = seed
1180-
self.units_seed = {unit_id: self.segment_seed + hash(unit_id) for unit_id in range(num_units)}
1182+
self.units_seed = {unit_id: abs(self.segment_seed + hash(unit_id)) for unit_id in unit_ids}
1183+
11811184
self.num_samples = math.ceil(sampling_frequency * duration)
11821185
super().__init__(t_start)
11831186

@@ -1280,7 +1283,7 @@ def __init__(
12801283
noise_block_size: int = 30000,
12811284
):
12821285

1283-
channel_ids = np.arange(num_channels)
1286+
channel_ids = [str(idx) for idx in np.arange(num_channels)]
12841287
dtype = np.dtype(dtype).name # Cast to string for serialization
12851288
if dtype not in ("float32", "float64"):
12861289
raise ValueError(f"'dtype' must be 'float32' or 'float64' but is {dtype}")

src/spikeinterface/core/tests/test_basesnippets.py

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -41,8 +41,8 @@ def test_BaseSnippets(create_cache_folder):
4141
assert snippets.get_num_segments() == len(duration)
4242
assert snippets.get_num_channels() == num_channels
4343

44-
assert np.all(snippets.ids_to_indices([0, 1, 2]) == [0, 1, 2])
45-
assert np.all(snippets.ids_to_indices([0, 1, 2], prefer_slice=True) == slice(0, 3, None))
44+
assert np.all(snippets.ids_to_indices(["0", "1", "2"]) == [0, 1, 2])
45+
assert np.all(snippets.ids_to_indices(["0", "1", "2"], prefer_slice=True) == slice(0, 3, None))
4646

4747
# annotations / properties
4848
snippets.annotate(gre="ta")
@@ -60,7 +60,7 @@ def test_BaseSnippets(create_cache_folder):
6060
)
6161

6262
# missing property
63-
snippets.set_property("string_property", ["ciao", "bello"], ids=[0, 1])
63+
snippets.set_property("string_property", ["ciao", "bello"], ids=["0", "1"])
6464
values = snippets.get_property("string_property")
6565
assert values[2] == ""
6666

@@ -70,14 +70,14 @@ def test_BaseSnippets(create_cache_folder):
7070
snippets.set_property,
7171
key="string_property_nan",
7272
values=["hola", "chabon"],
73-
ids=[0, 1],
73+
ids=["0", "1"],
7474
missing_value=np.nan,
7575
)
7676

7777
# int properties without missing values raise an error
7878
assert_raises(Exception, snippets.set_property, key="int_property", values=[5, 6], ids=[1, 2])
7979

80-
snippets.set_property("int_property", [5, 6], ids=[1, 2], missing_value=200)
80+
snippets.set_property("int_property", [5, 6], ids=["1", "2"], missing_value=200)
8181
values = snippets.get_property("int_property")
8282
assert values.dtype.kind == "i"
8383

src/spikeinterface/core/tests/test_channelsaggregationrecording.py

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -38,10 +38,12 @@ def test_channelsaggregationrecording():
3838

3939
assert np.allclose(traces1_1, recording_agg.get_traces(channel_ids=[str(channel_ids[1])], segment_index=seg))
4040
assert np.allclose(
41-
traces2_0, recording_agg.get_traces(channel_ids=[str(num_channels + channel_ids[0])], segment_index=seg)
41+
traces2_0,
42+
recording_agg.get_traces(channel_ids=[str(num_channels + int(channel_ids[0]))], segment_index=seg),
4243
)
4344
assert np.allclose(
44-
traces3_2, recording_agg.get_traces(channel_ids=[str(2 * num_channels + channel_ids[2])], segment_index=seg)
45+
traces3_2,
46+
recording_agg.get_traces(channel_ids=[str(2 * num_channels + int(channel_ids[2]))], segment_index=seg),
4547
)
4648
# all traces
4749
traces1 = recording1.get_traces(segment_index=seg)

src/spikeinterface/core/tests/test_sortinganalyzer.py

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -31,6 +31,14 @@ def get_dataset():
3131
noise_kwargs=dict(noise_levels=5.0, strategy="tile_pregenerated"),
3232
seed=2205,
3333
)
34+
35+
# TODO: the tests or the sorting analyzer make assumptions about the ids being integers
36+
# So keeping this the way it was
37+
integer_channel_ids = [int(id) for id in recording.get_channel_ids()]
38+
integer_unit_ids = [int(id) for id in sorting.get_unit_ids()]
39+
40+
recording = recording.rename_channels(new_channel_ids=integer_channel_ids)
41+
sorting = sorting.rename_units(new_unit_ids=integer_unit_ids)
3442
return recording, sorting
3543

3644

src/spikeinterface/core/tests/test_unitsselectionsorting.py

Lines changed: 13 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -10,39 +10,43 @@
1010
def test_basic_functions():
1111
sorting = generate_sorting(num_units=3, durations=[0.100, 0.100], sampling_frequency=30000.0)
1212

13-
sorting2 = UnitsSelectionSorting(sorting, unit_ids=[0, 2])
14-
assert np.array_equal(sorting2.unit_ids, [0, 2])
13+
sorting2 = UnitsSelectionSorting(sorting, unit_ids=["0", "2"])
14+
assert np.array_equal(sorting2.unit_ids, ["0", "2"])
1515
assert sorting2.get_parent() == sorting
1616

17-
sorting3 = UnitsSelectionSorting(sorting, unit_ids=[0, 2], renamed_unit_ids=["a", "b"])
17+
sorting3 = UnitsSelectionSorting(sorting, unit_ids=["0", "2"], renamed_unit_ids=["a", "b"])
1818
assert np.array_equal(sorting3.unit_ids, ["a", "b"])
1919

2020
assert np.array_equal(
21-
sorting.get_unit_spike_train(0, segment_index=0), sorting2.get_unit_spike_train(0, segment_index=0)
21+
sorting.get_unit_spike_train(unit_id="0", segment_index=0),
22+
sorting2.get_unit_spike_train(unit_id="0", segment_index=0),
2223
)
2324
assert np.array_equal(
24-
sorting.get_unit_spike_train(0, segment_index=0), sorting3.get_unit_spike_train("a", segment_index=0)
25+
sorting.get_unit_spike_train(unit_id="0", segment_index=0),
26+
sorting3.get_unit_spike_train(unit_id="a", segment_index=0),
2527
)
2628

2729
assert np.array_equal(
28-
sorting.get_unit_spike_train(2, segment_index=0), sorting2.get_unit_spike_train(2, segment_index=0)
30+
sorting.get_unit_spike_train(unit_id="2", segment_index=0),
31+
sorting2.get_unit_spike_train(unit_id="2", segment_index=0),
2932
)
3033
assert np.array_equal(
31-
sorting.get_unit_spike_train(2, segment_index=0), sorting3.get_unit_spike_train("b", segment_index=0)
34+
sorting.get_unit_spike_train(unit_id="2", segment_index=0),
35+
sorting3.get_unit_spike_train(unit_id="b", segment_index=0),
3236
)
3337

3438

3539
def test_failure_with_non_unique_unit_ids():
3640
seed = 10
3741
sorting = generate_sorting(num_units=3, durations=[0.100], sampling_frequency=30000.0, seed=seed)
3842
with pytest.raises(AssertionError):
39-
sorting2 = UnitsSelectionSorting(sorting, unit_ids=[0, 2], renamed_unit_ids=["a", "a"])
43+
sorting2 = UnitsSelectionSorting(sorting, unit_ids=["0", "2"], renamed_unit_ids=["a", "a"])
4044

4145

4246
def test_custom_cache_spike_vector():
4347
sorting = generate_sorting(num_units=3, durations=[0.100, 0.100], sampling_frequency=30000.0)
4448

45-
sub_sorting = UnitsSelectionSorting(sorting, unit_ids=[2, 0], renamed_unit_ids=["b", "a"])
49+
sub_sorting = UnitsSelectionSorting(sorting, unit_ids=["2", "0"], renamed_unit_ids=["b", "a"])
4650
cached_spike_vector = sub_sorting.to_spike_vector(use_cache=True)
4751
computed_spike_vector = sub_sorting.to_spike_vector(use_cache=False)
4852
assert np.all(cached_spike_vector == computed_spike_vector)

src/spikeinterface/curation/tests/common.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,11 @@ def make_sorting_analyzer(sparse=True):
1919
seed=2205,
2020
)
2121

22+
channel_ids_as_integers = [id for id in range(recording.get_num_channels())]
23+
unit_ids_as_integers = [id for id in range(sorting.get_num_units())]
24+
recording = recording.rename_channels(new_channel_ids=channel_ids_as_integers)
25+
sorting = sorting.rename_units(new_unit_ids=unit_ids_as_integers)
26+
2227
sorting_analyzer = create_sorting_analyzer(sorting=sorting, recording=recording, format="memory", sparse=sparse)
2328
sorting_analyzer.compute("random_spikes")
2429
sorting_analyzer.compute("waveforms", **job_kwargs)

src/spikeinterface/curation/tests/test_sortingview_curation.py

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -49,6 +49,9 @@ def test_gh_curation():
4949
Test curation using GitHub URI.
5050
"""
5151
sorting = generate_sorting(num_units=10)
52+
unit_ids_as_int = [id for id in range(sorting.get_num_units())]
53+
sorting = sorting.rename_units(new_unit_ids=unit_ids_as_int)
54+
5255
# curated link:
5356
# https://figurl.org/f?v=npm://@fi-sci/figurl-sortingview@12/dist&d=sha1://058ab901610aa9d29df565595a3cc2a81a1b08e5
5457
gh_uri = "gh://SpikeInterface/spikeinterface/main/src/spikeinterface/curation/tests/sv-sorting-curation.json"
@@ -76,6 +79,8 @@ def test_sha1_curation():
7679
Test curation using SHA1 URI.
7780
"""
7881
sorting = generate_sorting(num_units=10)
82+
unit_ids_as_int = [id for id in range(sorting.get_num_units())]
83+
sorting = sorting.rename_units(new_unit_ids=unit_ids_as_int)
7984

8085
# from SHA1
8186
# curated link:
@@ -105,6 +110,8 @@ def test_json_curation():
105110
Test curation using a JSON file.
106111
"""
107112
sorting = generate_sorting(num_units=10)
113+
unit_ids_as_int = [id for id in range(sorting.get_num_units())]
114+
sorting = sorting.rename_units(new_unit_ids=unit_ids_as_int)
108115

109116
# from curation.json
110117
json_file = parent_folder / "sv-sorting-curation.json"
@@ -248,6 +255,8 @@ def test_json_no_merge_curation():
248255
Test curation with no merges using a JSON file.
249256
"""
250257
sorting = generate_sorting(num_units=10)
258+
unit_ids_as_int = [id for id in range(sorting.get_num_units())]
259+
sorting = sorting.rename_units(new_unit_ids=unit_ids_as_int)
251260

252261
json_file = parent_folder / "sv-sorting-curation-no-merge.json"
253262
sorting_curated = apply_sortingview_curation(sorting, uri_or_json=json_file)

src/spikeinterface/extractors/tests/test_mdaextractors.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,12 @@ def test_mda_extractors(create_cache_folder):
99
cache_folder = create_cache_folder
1010
rec, sort = generate_ground_truth_recording(durations=[10.0], num_units=10)
1111

12+
ids_as_integers = [id for id in range(rec.get_num_channels())]
13+
rec = rec.rename_channels(new_channel_ids=ids_as_integers)
14+
15+
ids_as_integers = [id for id in range(sort.get_num_units())]
16+
sort = sort.rename_units(new_unit_ids=ids_as_integers)
17+
1218
MdaRecordingExtractor.write_recording(rec, cache_folder / "mdatest")
1319
rec_mda = MdaRecordingExtractor(cache_folder / "mdatest")
1420
probe = rec_mda.get_probe()

src/spikeinterface/postprocessing/tests/test_multi_extensions.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,11 @@ def get_dataset():
2323
seed=2205,
2424
)
2525

26+
channel_ids_as_integers = [id for id in range(recording.get_num_channels())]
27+
unit_ids_as_integers = [id for id in range(sorting.get_num_units())]
28+
recording = recording.rename_channels(new_channel_ids=channel_ids_as_integers)
29+
sorting = sorting.rename_units(new_unit_ids=unit_ids_as_integers)
30+
2631
# since templates are going to be averaged and this might be a problem for amplitude scaling
2732
# we select the 3 units with the largest templates to split
2833
analyzer_raw = create_sorting_analyzer(sorting, recording, format="memory", sparse=False)

0 commit comments

Comments
 (0)