diff --git a/docs/source/_static/interactive-multi-example.png b/docs/source/_static/interactive-multi-example.png index 31ceb10..a2ae137 100644 Binary files a/docs/source/_static/interactive-multi-example.png and b/docs/source/_static/interactive-multi-example.png differ diff --git a/docs/source/_static/interactive-single-example.png b/docs/source/_static/interactive-single-example.png index c42b8fb..7c3820d 100644 Binary files a/docs/source/_static/interactive-single-example.png and b/docs/source/_static/interactive-single-example.png differ diff --git a/docs/source/pages/how-to-use-driftplots.md b/docs/source/pages/how-to-use-driftplots.md index f1757eb..66d40a7 100644 --- a/docs/source/pages/how-to-use-driftplots.md +++ b/docs/source/pages/how-to-use-driftplots.md @@ -28,8 +28,7 @@ and not reflect any later changes in Phy (`spike_templates.npy` is used for the If passing a `SortingAnalyzer`, it is expected that the required extensions have already been computed. See [this example](https://github.com/neuroinformatics-unit/driftplots/blob/e8ec328e14cc848feca3e7e90604501bb9e343f1/examples/example_data/create_analyzer.py#L1) -for the required extensions. Note that the number of spikes displayed will depend -on the argument set for `max_spikes_per_unit` used when computing `"random_spikes"`. +for the required extensions. By default, the number of spikes displayed will be decimated to around `100,000`. @@ -194,5 +193,3 @@ for path_or_analyzer in SORTING_SESSIONS: multi = MultiSessionDriftmapWidget(panels) multi.plot() - -``` diff --git a/driftplots/extractors/analyzer_helpers.py b/driftplots/extractors/analyzer_helpers.py index 285eb9a..a615278 100644 --- a/driftplots/extractors/analyzer_helpers.py +++ b/driftplots/extractors/analyzer_helpers.py @@ -10,26 +10,19 @@ def get_sorting_analyzer( analyzer: si.SortingAnalyzer, ) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray, np.ndarray, np.ndarray]: """ - Get the required data from the SortingAnalyzer. Note that this - will not get all detected spikes, but rather the number of spikes - specified when creating the analyzer, `max_spikes_per_unit`. + Get the required data from the SortingAnalyzer. """ - random_spike_indices = analyzer.get_extension("random_spikes").data[ - "random_spikes_indices" - ] spike_vector = analyzer.sorting.to_spike_vector() spike_times = ( - spike_vector["sample_index"][random_spike_indices] - / analyzer.sorting.get_sampling_frequency() + spike_vector["sample_index"] / analyzer.sorting.get_sampling_frequency() ) spike_amplitudes = np.abs( analyzer.get_extension("spike_amplitudes").data["amplitudes"] ) - spike_depths = analyzer.get_extension("spike_locations").data["spike_locations"][ "y" ] - spike_templates = spike_vector["unit_index"][random_spike_indices] + spike_templates = spike_vector["unit_index"] # Get the templates, assume only one method was used. If multiple # methods were used, use the first and throw a warning. If people diff --git a/examples/example_data/create_analyzer.py b/examples/example_data/create_analyzer.py index efd57aa..640b11b 100644 --- a/examples/example_data/create_analyzer.py +++ b/examples/example_data/create_analyzer.py @@ -17,9 +17,9 @@ rec = si_prepro.common_reference(rec, operator="median") if SAVE_SORTING: - # out_path = base_path / "sorting" - # out_path.mkdir() - sort = run_sorter("kilosort4", rec, folder=base_path / "sorting") + out_path = base_path / "sorting" + out_path.mkdir() + sort = run_sorter("kilosort4", rec, folder=out_path) else: sort = si_extractors.read_kilosort( base_path / "sorting" / "kilosort4_output" / "sorter_output" @@ -29,9 +29,6 @@ analyzer.compute( "random_spikes", method="uniform", - # This determines the number of spikes that - # will appear on the SI drift plot - max_spikes_per_unit=1_000_000, ) analyzer.compute("waveforms", ms_before=1.0, ms_after=2.0) analyzer.compute("templates", operators=["average"]) diff --git a/tests/test_sorting_analyzer.py b/tests/test_sorting_analyzer.py index 61a93dd..0d3f7e5 100644 --- a/tests/test_sorting_analyzer.py +++ b/tests/test_sorting_analyzer.py @@ -22,15 +22,10 @@ def test_loaded_arrays_match_analyzer_extensions(self): analyzer = si.load_sorting_analyzer(ANALYZER_PATH) loader = DataLoader(analyzer, verbose=False) - # Expected values from the analyzer extensions - random_spike_indices = analyzer.get_extension("random_spikes").data[ - "random_spikes_indices" - ] spike_vector = analyzer.sorting.to_spike_vector() expected_times = ( - spike_vector["sample_index"][random_spike_indices] - / analyzer.sorting.get_sampling_frequency() + spike_vector["sample_index"] / analyzer.sorting.get_sampling_frequency() ) np.testing.assert_array_equal(loader._spike_times, expected_times) @@ -44,7 +39,7 @@ def test_loaded_arrays_match_analyzer_extensions(self): ]["y"] np.testing.assert_array_equal(loader._spike_depths, expected_depths) - expected_templates = spike_vector["unit_index"][random_spike_indices] + expected_templates = spike_vector["unit_index"] np.testing.assert_array_equal(loader._spike_templates, expected_templates) expected_waveforms = analyzer.get_extension("templates").data["average"] @@ -65,6 +60,27 @@ def test_loaded_arrays_match_analyzer_extensions(self): ): assert not arr.flags.writeable + def test_loaded_spikes_ignore_random_spikes_subset(self): + analyzer = si.load_sorting_analyzer(ANALYZER_PATH).copy() + spike_vector = analyzer.sorting.to_spike_vector() + + random_spikes = analyzer.get_extension("random_spikes") + random_spikes.data["random_spikes_indices"] = random_spikes.data[ + "random_spikes_indices" + ][::2] + assert random_spikes.data["random_spikes_indices"].size < spike_vector.size + + loader = DataLoader(analyzer, verbose=False) + + expected_times = ( + spike_vector["sample_index"] / analyzer.sorting.get_sampling_frequency() + ) + np.testing.assert_array_equal(loader._spike_times, expected_times) + np.testing.assert_array_equal( + loader._spike_templates, spike_vector["unit_index"] + ) + assert loader._spike_times.size == spike_vector.size + def test_good_units_only_keeps_kslabel_good_units(self): analyzer = si.load_sorting_analyzer(ANALYZER_PATH) loader = DataLoader(analyzer, verbose=False)