|
10 | 10 | def test_basic_functions(): |
11 | 11 | sorting = generate_sorting(num_units=3, durations=[0.100, 0.100], sampling_frequency=30000.0) |
12 | 12 |
|
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"]) |
15 | 15 | assert sorting2.get_parent() == sorting |
16 | 16 |
|
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"]) |
18 | 18 | assert np.array_equal(sorting3.unit_ids, ["a", "b"]) |
19 | 19 |
|
20 | 20 | 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), |
22 | 23 | ) |
23 | 24 | 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), |
25 | 27 | ) |
26 | 28 |
|
27 | 29 | 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), |
29 | 32 | ) |
30 | 33 | 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), |
32 | 36 | ) |
33 | 37 |
|
34 | 38 |
|
35 | 39 | def test_failure_with_non_unique_unit_ids(): |
36 | 40 | seed = 10 |
37 | 41 | sorting = generate_sorting(num_units=3, durations=[0.100], sampling_frequency=30000.0, seed=seed) |
38 | 42 | 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"]) |
40 | 44 |
|
41 | 45 |
|
42 | 46 | def test_custom_cache_spike_vector(): |
43 | 47 | sorting = generate_sorting(num_units=3, durations=[0.100, 0.100], sampling_frequency=30000.0) |
44 | 48 |
|
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"]) |
46 | 50 | cached_spike_vector = sub_sorting.to_spike_vector(use_cache=True) |
47 | 51 | computed_spike_vector = sub_sorting.to_spike_vector(use_cache=False) |
48 | 52 | assert np.all(cached_spike_vector == computed_spike_vector) |
|
0 commit comments