Skip to content

Commit b6316b5

Browse files
authored
Merge branch 'SpikeInterface:main' into sortingview-curation-fix
2 parents d029f7d + 694f862 commit b6316b5

1 file changed

Lines changed: 31 additions & 3 deletions

File tree

src/spikeinterface/sorters/external/kilosort4.py

Lines changed: 31 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -56,6 +56,7 @@ class Kilosort4Sorter(BaseSorter):
5656
"save_extra_kwargs": False,
5757
"skip_kilosort_preprocessing": False,
5858
"scaleproc": None,
59+
"save_preprocessed_copy": False,
5960
"torch_device": "auto",
6061
}
6162

@@ -98,6 +99,7 @@ class Kilosort4Sorter(BaseSorter):
9899
"save_extra_kwargs": "If True, additional kwargs are saved to the output",
99100
"skip_kilosort_preprocessing": "Can optionally skip the internal kilosort preprocessing",
100101
"scaleproc": "int16 scaling of whitened data, if None set to 200.",
102+
"save_preprocessed_copy": "save a pre-processed copy of the data (including drift correction) to temp_wh.dat in the results directory and format Phy output to use that copy of the data",
101103
"torch_device": "Select the torch device auto/cuda/cpu",
102104
}
103105

@@ -153,7 +155,7 @@ def _run_from_folder(cls, sorter_output_folder, params, verbose):
153155
save_sorting,
154156
get_run_parameters,
155157
)
156-
from kilosort.io import load_probe, RecordingExtractorAsArray, BinaryFiltered
158+
from kilosort.io import load_probe, RecordingExtractorAsArray, BinaryFiltered, save_preprocessing
157159
from kilosort.parameters import DEFAULT_SETTINGS
158160

159161
import time
@@ -186,6 +188,7 @@ def _run_from_folder(cls, sorter_output_folder, params, verbose):
186188
do_CAR = params["do_CAR"]
187189
invert_sign = params["invert_sign"]
188190
save_extra_vars = params["save_extra_kwargs"]
191+
save_preprocessed_copy = params["save_preprocessed_copy"]
189192
progress_bar = None
190193
settings_ks = {k: v for k, v in params.items() if k in DEFAULT_SETTINGS}
191194
settings_ks["n_chan_bin"] = recording.get_num_channels()
@@ -207,7 +210,15 @@ def _run_from_folder(cls, sorter_output_folder, params, verbose):
207210
results_dir = sorter_output_folder
208211
filename, data_dir, results_dir, probe = set_files(settings, filename, probe, probe_name, data_dir, results_dir)
209212
if version.parse(cls.get_sorter_version()) >= version.parse("4.0.12"):
210-
ops = initialize_ops(settings, probe, recording.get_dtype(), do_CAR, invert_sign, device, False)
213+
ops = initialize_ops(
214+
settings,
215+
probe,
216+
recording.get_dtype(),
217+
do_CAR,
218+
invert_sign,
219+
device,
220+
save_preprocesed_copy=save_preprocessed_copy, # this kwarg is correct (typo)
221+
)
211222
n_chan_bin, fs, NT, nt, twav_min, chan_map, dtype, do_CAR, invert, _, _, tmin, tmax, artifact, _, _ = (
212223
get_run_parameters(ops)
213224
)
@@ -257,6 +268,9 @@ def _run_from_folder(cls, sorter_output_folder, params, verbose):
257268
ops, device, tic0=tic0, progress_bar=progress_bar, file_object=file_object
258269
)
259270

271+
if save_preprocessed_copy:
272+
save_preprocessing(results_dir / "temp_wh.dat", ops, bfile)
273+
260274
# Sort spikes and save results
261275
st, tF, _, _ = detect_spikes(ops, device, bfile, tic0=tic0, progress_bar=progress_bar)
262276
clu, Wall = cluster_spikes(st, tF, ops, device, bfile, tic0=tic0, progress_bar=progress_bar)
@@ -265,7 +279,21 @@ def _run_from_folder(cls, sorter_output_folder, params, verbose):
265279
hp_filter=torch.as_tensor(np.zeros(1)), whiten_mat=torch.as_tensor(np.eye(recording.get_num_channels()))
266280
)
267281

268-
_ = save_sorting(ops, results_dir, st, clu, tF, Wall, bfile.imin, tic0, save_extra_vars=save_extra_vars)
282+
if version.parse(cls.get_sorter_version()) >= version.parse("4.0.12"):
283+
_ = save_sorting(
284+
ops,
285+
results_dir,
286+
st,
287+
clu,
288+
tF,
289+
Wall,
290+
bfile.imin,
291+
tic0,
292+
save_extra_vars=save_extra_vars,
293+
save_preprocessed_copy=save_preprocessed_copy,
294+
)
295+
else:
296+
_ = save_sorting(ops, results_dir, st, clu, tF, Wall, bfile.imin, tic0, save_extra_vars=save_extra_vars)
269297

270298
@classmethod
271299
def _get_result_from_folder(cls, sorter_output_folder):

0 commit comments

Comments
 (0)