Skip to content

Commit bac57fe

Browse files
[pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
1 parent 52e8d2a commit bac57fe

7 files changed

Lines changed: 54 additions & 42 deletions

File tree

src/spikeinterface/core/analyzer_extension_core.py

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -693,7 +693,6 @@ class ComputeNoiseLevels(AnalyzerExtension):
693693
need_job_kwargs = False
694694
need_backward_compatibility_on_load = True
695695

696-
697696
def __init__(self, sorting_analyzer):
698697
AnalyzerExtension.__init__(self, sorting_analyzer)
699698

src/spikeinterface/core/job_tools.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -184,6 +184,7 @@ def ensure_n_jobs(recording, n_jobs=1):
184184

185185
return n_jobs
186186

187+
187188
def chunk_duration_to_chunk_size(chunk_duration, recording):
188189
if isinstance(chunk_duration, float):
189190
chunk_size = int(chunk_duration * recording.get_sampling_frequency())
@@ -196,7 +197,7 @@ def chunk_duration_to_chunk_size(chunk_duration, recording):
196197
raise ValueError("chunk_duration must ends with s or ms")
197198
chunk_size = int(chunk_duration * recording.get_sampling_frequency())
198199
else:
199-
raise ValueError("chunk_duration must be str or float")
200+
raise ValueError("chunk_duration must be str or float")
200201
return chunk_size
201202

202203

src/spikeinterface/core/recording_tools.py

Lines changed: 27 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -514,15 +514,15 @@ def determine_cast_unsigned(recording, dtype):
514514
return cast_unsigned
515515

516516

517-
518-
519-
def get_random_recording_slices(recording,
520-
method="full_random",
521-
num_chunks_per_segment=20,
522-
chunk_duration="500ms",
523-
chunk_size=None,
524-
margin_frames=0,
525-
seed=None):
517+
def get_random_recording_slices(
518+
recording,
519+
method="full_random",
520+
num_chunks_per_segment=20,
521+
chunk_duration="500ms",
522+
chunk_size=None,
523+
margin_frames=0,
524+
seed=None,
525+
):
526526
"""
527527
Get random slice of a recording across segments.
528528
@@ -593,19 +593,14 @@ def get_random_recording_slices(recording,
593593
]
594594
else:
595595
raise ValueError(f"get_random_recording_slices : wrong method {method}")
596-
596+
597597
return recording_slices
598598

599599

600-
def get_random_data_chunks(
601-
recording,
602-
return_scaled=False,
603-
concatenated=True,
604-
**random_slices_kwargs
605-
):
600+
def get_random_data_chunks(recording, return_scaled=False, concatenated=True, **random_slices_kwargs):
606601
"""
607602
Extract random chunks across segments.
608-
603+
609604
Internally, it uses `get_random_recording_slices()` and retrieves the traces chunk as a list
610605
or a concatenated unique array.
611606
@@ -698,15 +693,14 @@ def get_closest_channels(recording, channel_ids=None, num_channels=None):
698693

699694
def _noise_level_chunk(segment_index, start_frame, end_frame, worker_ctx):
700695
recording = worker_ctx["recording"]
701-
696+
702697
one_chunk = recording.get_traces(
703698
start_frame=start_frame,
704699
end_frame=end_frame,
705700
segment_index=segment_index,
706701
return_scaled=worker_ctx["return_scaled"],
707702
)
708703

709-
710704
if worker_ctx["method"] == "mad":
711705
med = np.median(one_chunk, axis=0, keepdims=True)
712706
# hard-coded so that core doesn't depend on scipy
@@ -724,12 +718,13 @@ def _noise_level_chunk_init(recording, return_scaled, method):
724718
worker_ctx["method"] = method
725719
return worker_ctx
726720

721+
727722
def get_noise_levels(
728723
recording: "BaseRecording",
729724
return_scaled: bool = True,
730725
method: Literal["mad", "std"] = "mad",
731726
force_recompute: bool = False,
732-
random_slices_kwargs : dict = {},
727+
random_slices_kwargs: dict = {},
733728
**kwargs,
734729
) -> np.ndarray:
735730
"""
@@ -759,7 +754,7 @@ def get_noise_levels(
759754
function for more details.
760755
761756
{}
762-
757+
763758
Returns
764759
-------
765760
noise_levels : array
@@ -774,7 +769,7 @@ def get_noise_levels(
774769
if key in recording.get_property_keys() and not force_recompute:
775770
noise_levels = recording.get_property(key=key)
776771
else:
777-
# This is to keep backward compatibility
772+
# This is to keep backward compatibility
778773
# lets keep for a while and remove this maybe in 0.103.0
779774
# chunk_size used to be in the signature and now is ambiguous
780775
random_slices_kwargs_, job_kwargs = split_job_kwargs(kwargs)
@@ -794,15 +789,22 @@ def get_noise_levels(
794789
recording_slices = get_random_recording_slices(recording, **random_slices_kwargs)
795790

796791
noise_levels_chunks = []
792+
797793
def append_noise_chunk(res):
798794
noise_levels_chunks.append(res)
799795

800796
func = _noise_level_chunk
801797
init_func = _noise_level_chunk_init
802798
init_args = (recording, return_scaled, method)
803799
executor = ChunkRecordingExecutor(
804-
recording, func, init_func, init_args, job_name="noise_level", verbose=False,
805-
gather_func=append_noise_chunk, **job_kwargs
800+
recording,
801+
func,
802+
init_func,
803+
init_args,
804+
job_name="noise_level",
805+
verbose=False,
806+
gather_func=append_noise_chunk,
807+
**job_kwargs,
806808
)
807809
executor.run(all_chunks=recording_slices)
808810
noise_levels_chunks = np.stack(noise_levels_chunks)
@@ -813,6 +815,7 @@ def append_noise_chunk(res):
813815

814816
return noise_levels
815817

818+
816819
get_noise_levels.__doc__ = get_noise_levels.__doc__.format(_shared_job_kwargs_doc)
817820

818821

src/spikeinterface/core/tests/test_analyzer_extension_core.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -259,7 +259,7 @@ def test_compute_several(create_cache_folder):
259259
# test_ComputeWaveforms(format="binary_folder", sparse=False, create_cache_folder=cache_folder)
260260
# test_ComputeWaveforms(format="zarr", sparse=True, create_cache_folder=cache_folder)
261261
# test_ComputeWaveforms(format="zarr", sparse=False, create_cache_folder=cache_folder)
262-
#test_ComputeRandomSpikes(format="memory", sparse=True, create_cache_folder=cache_folder)
262+
# test_ComputeRandomSpikes(format="memory", sparse=True, create_cache_folder=cache_folder)
263263
test_ComputeRandomSpikes(format="binary_folder", sparse=False, create_cache_folder=cache_folder)
264264
test_ComputeTemplates(format="memory", sparse=True, create_cache_folder=cache_folder)
265265
test_ComputeNoiseLevels(format="memory", sparse=False, create_cache_folder=cache_folder)

src/spikeinterface/core/tests/test_recording_tools.py

Lines changed: 22 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -167,19 +167,18 @@ def test_write_memory_recording():
167167
for shm in shms:
168168
shm.unlink()
169169

170+
170171
def test_get_random_recording_slices():
171172
rec = generate_recording(num_channels=1, sampling_frequency=1000.0, durations=[10.0, 20.0])
172-
rec_slices = get_random_recording_slices(rec,
173-
method="full_random",
174-
num_chunks_per_segment=20,
175-
chunk_duration="500ms",
176-
margin_frames=0,
177-
seed=0)
173+
rec_slices = get_random_recording_slices(
174+
rec, method="full_random", num_chunks_per_segment=20, chunk_duration="500ms", margin_frames=0, seed=0
175+
)
178176
assert len(rec_slices) == 40
179177
for seg_ind, start, stop in rec_slices:
180178
assert stop - start == 500
181179
assert seg_ind in (0, 1)
182180

181+
183182
def test_get_random_data_chunks():
184183
rec = generate_recording(num_channels=1, sampling_frequency=1000.0, durations=[10.0, 20.0])
185184
chunks = get_random_data_chunks(rec, num_chunks_per_segment=50, chunk_size=500, seed=0)
@@ -216,7 +215,9 @@ def test_get_noise_levels():
216215

217216
assert np.all(noise_levels_1 == noise_levels_2)
218217
assert np.allclose(get_noise_levels(recording, return_scaled=False, **job_kwargs), [std, std], rtol=1e-2, atol=1e-3)
219-
assert np.allclose(get_noise_levels(recording, method="std", return_scaled=False, **job_kwargs), [std, std], rtol=1e-2, atol=1e-3)
218+
assert np.allclose(
219+
get_noise_levels(recording, method="std", return_scaled=False, **job_kwargs), [std, std], rtol=1e-2, atol=1e-3
220+
)
220221

221222

222223
def test_get_noise_levels_output():
@@ -230,13 +231,21 @@ def test_get_noise_levels_output():
230231
traces = rng.normal(loc=10.0, scale=std, size=(num_samples, num_channels))
231232
recording = NumpyRecording(traces_list=traces, sampling_frequency=sampling_frequency)
232233

233-
std_estimated_with_mad = get_noise_levels(recording, method="mad", return_scaled=False,
234-
random_slices_kwargs=dict(num_chunks_per_segment=40, chunk_size=1_000, seed=seed))
234+
std_estimated_with_mad = get_noise_levels(
235+
recording,
236+
method="mad",
237+
return_scaled=False,
238+
random_slices_kwargs=dict(num_chunks_per_segment=40, chunk_size=1_000, seed=seed),
239+
)
235240
print(std_estimated_with_mad)
236241
assert np.allclose(std_estimated_with_mad, [std, std], rtol=1e-2, atol=1e-3)
237242

238-
std_estimated_with_std = get_noise_levels(recording, method="std", return_scaled=False,
239-
random_slices_kwargs=dict(num_chunks_per_segment=40, chunk_size=1_000, seed=seed))
243+
std_estimated_with_std = get_noise_levels(
244+
recording,
245+
method="std",
246+
return_scaled=False,
247+
random_slices_kwargs=dict(num_chunks_per_segment=40, chunk_size=1_000, seed=seed),
248+
)
240249
assert np.allclose(std_estimated_with_std, [std, std], rtol=1e-2, atol=1e-3)
241250

242251

@@ -358,8 +367,8 @@ def test_do_recording_attributes_match():
358367
# test_write_memory_recording()
359368

360369
test_get_random_recording_slices()
361-
# test_get_random_data_chunks()
370+
# test_get_random_data_chunks()
362371
# test_get_closest_channels()
363372
# test_get_noise_levels()
364-
# test_get_noise_levels_output()
373+
# test_get_noise_levels_output()
365374
# test_order_channels_by_depth()

src/spikeinterface/preprocessing/tests/test_silence.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -48,5 +48,5 @@ def test_silence(create_cache_folder):
4848

4949

5050
if __name__ == "__main__":
51-
cache_folder = Path(__file__).resolve().parents[4] / "cache_folder"
51+
cache_folder = Path(__file__).resolve().parents[4] / "cache_folder"
5252
test_silence(cache_folder)

src/spikeinterface/preprocessing/tests/test_whiten.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -49,5 +49,5 @@ def test_whiten(create_cache_folder):
4949

5050

5151
if __name__ == "__main__":
52-
cache_folder = Path(__file__).resolve().parents[4] / "cache_folder"
52+
cache_folder = Path(__file__).resolve().parents[4] / "cache_folder"
5353
test_whiten(cache_folder)

0 commit comments

Comments
 (0)