Skip to content

Commit 469187a

Browse files
authored
Merge pull request #3347 from jonahpearl/analyzer_extension_exit_status
Analyzer extension exit status
2 parents 096d91a + cc21f06 commit 469187a

2 files changed

Lines changed: 116 additions & 7 deletions

File tree

src/spikeinterface/core/sortinganalyzer.py

Lines changed: 83 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,7 @@
1111
import shutil
1212
import warnings
1313
import importlib
14+
from time import perf_counter
1415

1516
import numpy as np
1617

@@ -1336,6 +1337,7 @@ def compute_several_extensions(self, extensions, save=True, verbose=False, **job
13361337

13371338
job_name = "Compute : " + " + ".join(extensions_with_pipeline.keys())
13381339

1340+
t_start = perf_counter()
13391341
results = run_node_pipeline(
13401342
self.recording,
13411343
all_nodes,
@@ -1345,10 +1347,15 @@ def compute_several_extensions(self, extensions, save=True, verbose=False, **job
13451347
squeeze_output=False,
13461348
verbose=verbose,
13471349
)
1350+
t_end = perf_counter()
1351+
# for pipeline node extensions we can only track the runtime of the run_node_pipeline
1352+
runtime_s = t_end - t_start
13481353

13491354
for r, result in enumerate(results):
13501355
extension_name, variable_name = result_routage[r]
13511356
extension_instances[extension_name].data[variable_name] = result
1357+
extension_instances[extension_name].run_info["runtime_s"] = runtime_s
1358+
extension_instances[extension_name].run_info["run_completed"] = True
13521359

13531360
for extension_name, extension_instance in extension_instances.items():
13541361
self.extensions[extension_name] = extension_instance
@@ -1738,8 +1745,12 @@ def __init__(self, sorting_analyzer):
17381745
self._sorting_analyzer = weakref.ref(sorting_analyzer)
17391746

17401747
self.params = None
1748+
self.run_info = self._default_run_info_dict()
17411749
self.data = dict()
17421750

1751+
def _default_run_info_dict(self):
1752+
return dict(run_completed=False, runtime_s=None)
1753+
17431754
#######
17441755
# This 3 methods must be implemented in the subclass!!!
17451756
# See DummyAnalyzerExtension in test_sortinganalyzer.py as a simple example
@@ -1851,11 +1862,42 @@ def _get_zarr_extension_group(self, mode="r+"):
18511862
def load(cls, sorting_analyzer):
18521863
ext = cls(sorting_analyzer)
18531864
ext.load_params()
1854-
ext.load_data()
1855-
if cls.need_backward_compatibility_on_load:
1856-
ext._handle_backward_compatibility_on_load()
1865+
ext.load_run_info()
1866+
if ext.run_info is not None:
1867+
if ext.run_info["run_completed"]:
1868+
ext.load_data()
1869+
if cls.need_backward_compatibility_on_load:
1870+
ext._handle_backward_compatibility_on_load()
1871+
if len(ext.data) > 0:
1872+
return ext
1873+
else:
1874+
# this is for back-compatibility of old analyzers
1875+
ext.load_data()
1876+
if cls.need_backward_compatibility_on_load:
1877+
ext._handle_backward_compatibility_on_load()
1878+
if len(ext.data) > 0:
1879+
return ext
1880+
# If extension run not completed, or data has gone missing,
1881+
# return None to indicate that the extension should be (re)computed.
1882+
return None
1883+
1884+
def load_run_info(self):
1885+
if self.format == "binary_folder":
1886+
extension_folder = self._get_binary_extension_folder()
1887+
run_info_file = extension_folder / "run_info.json"
1888+
if run_info_file.is_file():
1889+
with open(str(run_info_file), "r") as f:
1890+
run_info = json.load(f)
1891+
else:
1892+
warnings.warn(f"Found no run_info file for {self.extension_name}, extension should be re-computed.")
1893+
run_info = None
18571894

1858-
return ext
1895+
elif self.format == "zarr":
1896+
extension_group = self._get_zarr_extension_group(mode="r")
1897+
run_info = extension_group.attrs.get("run_info", None)
1898+
if run_info is None:
1899+
warnings.warn(f"Found no run_info file for {self.extension_name}, extension should be re-computed.")
1900+
self.run_info = run_info
18591901

18601902
def load_params(self):
18611903
if self.format == "binary_folder":
@@ -1873,12 +1915,17 @@ def load_params(self):
18731915
self.params = params
18741916

18751917
def load_data(self):
1918+
ext_data = None
18761919
if self.format == "binary_folder":
18771920
extension_folder = self._get_binary_extension_folder()
18781921
for ext_data_file in extension_folder.iterdir():
18791922
# patch for https://github.com/SpikeInterface/spikeinterface/issues/3041
18801923
# maybe add a check for version number from the info.json during loading only
1881-
if ext_data_file.name == "params.json" or ext_data_file.name == "info.json":
1924+
if (
1925+
ext_data_file.name == "params.json"
1926+
or ext_data_file.name == "info.json"
1927+
or ext_data_file.name == "run_info.json"
1928+
):
18821929
continue
18831930
ext_data_name = ext_data_file.stem
18841931
if ext_data_file.suffix == ".json":
@@ -1919,6 +1966,9 @@ def load_data(self):
19191966
ext_data = np.array(ext_data_)
19201967
self.data[ext_data_name] = ext_data
19211968

1969+
if len(self.data) == 0:
1970+
warnings.warn(f"Found no data for {self.extension_name}, extension should be re-computed.")
1971+
19221972
def copy(self, new_sorting_analyzer, unit_ids=None):
19231973
# alessio : please note that this also replace the old select_units!!!
19241974
new_extension = self.__class__(new_sorting_analyzer)
@@ -1927,6 +1977,7 @@ def copy(self, new_sorting_analyzer, unit_ids=None):
19271977
new_extension.data = self.data
19281978
else:
19291979
new_extension.data = self._select_extension_data(unit_ids)
1980+
new_extension.run_info = self.run_info.copy()
19301981
new_extension.save()
19311982
return new_extension
19321983

@@ -1944,24 +1995,33 @@ def merge(
19441995
new_extension.data = self._merge_extension_data(
19451996
merge_unit_groups, new_unit_ids, new_sorting_analyzer, keep_mask, verbose=verbose, **job_kwargs
19461997
)
1998+
new_extension.run_info = self.run_info.copy()
19471999
new_extension.save()
19482000
return new_extension
19492001

19502002
def run(self, save=True, **kwargs):
19512003
if save and not self.sorting_analyzer.is_read_only():
1952-
# this also reset the folder or zarr group
2004+
# NB: this call to _save_params() also resets the folder or zarr group
19532005
self._save_params()
19542006
self._save_importing_provenance()
2007+
self._save_run_info()
19552008

2009+
t_start = perf_counter()
19562010
self._run(**kwargs)
2011+
t_end = perf_counter()
2012+
self.run_info["runtime_s"] = t_end - t_start
19572013

19582014
if save and not self.sorting_analyzer.is_read_only():
19592015
self._save_data(**kwargs)
19602016

2017+
self.run_info["run_completed"] = True
2018+
self._save_run_info()
2019+
19612020
def save(self, **kwargs):
19622021
self._save_params()
19632022
self._save_importing_provenance()
19642023
self._save_data(**kwargs)
2024+
self._save_run_info()
19652025

19662026
def _save_data(self, **kwargs):
19672027
if self.format == "memory":
@@ -2060,6 +2120,7 @@ def reset(self):
20602120
"""
20612121
self._reset_extension_folder()
20622122
self.params = None
2123+
self.run_info = self._default_run_info_dict()
20632124
self.data = dict()
20642125

20652126
def set_params(self, save=True, **params):
@@ -2080,6 +2141,7 @@ def set_params(self, save=True, **params):
20802141
if save:
20812142
self._save_params()
20822143
self._save_importing_provenance()
2144+
self._save_run_info()
20832145

20842146
def _save_params(self):
20852147
params_to_save = self.params.copy()
@@ -2117,14 +2179,28 @@ def _save_importing_provenance(self):
21172179
extension_group = self._get_zarr_extension_group(mode="r+")
21182180
extension_group.attrs["info"] = info
21192181

2182+
def _save_run_info(self):
2183+
run_info = self.run_info.copy()
2184+
2185+
if self.format == "binary_folder":
2186+
extension_folder = self._get_binary_extension_folder()
2187+
run_info_file = extension_folder / "run_info.json"
2188+
run_info_file.write_text(json.dumps(run_info, indent=4), encoding="utf8")
2189+
elif self.format == "zarr":
2190+
extension_group = self._get_zarr_extension_group(mode="r+")
2191+
extension_group.attrs["run_info"] = run_info
2192+
21202193
def get_pipeline_nodes(self):
21212194
assert (
21222195
self.use_nodepipeline
21232196
), "AnalyzerExtension.get_pipeline_nodes() must be called only when use_nodepipeline=True"
21242197
return self._get_pipeline_nodes()
21252198

21262199
def get_data(self, *args, **kwargs):
2127-
assert len(self.data) > 0, f"You must run the extension {self.extension_name} before retrieving data"
2200+
assert self.run_info[
2201+
"run_completed"
2202+
], f"You must run the extension {self.extension_name} before retrieving data"
2203+
assert len(self.data) > 0, "Extension has been run but no data found."
21282204
return self._get_data(*args, **kwargs)
21292205

21302206

src/spikeinterface/core/tests/test_sortinganalyzer.py

Lines changed: 33 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -125,6 +125,39 @@ def test_SortingAnalyzer_zarr(tmp_path, dataset):
125125
)
126126

127127

128+
def test_load_without_runtime_info(tmp_path, dataset):
129+
recording, sorting = dataset
130+
131+
folder = tmp_path / "test_SortingAnalyzer_run_info"
132+
133+
extensions = ["random_spikes", "templates"]
134+
# binary_folder
135+
sorting_analyzer = create_sorting_analyzer(
136+
sorting, recording, format="binary_folder", folder=folder, sparse=False, sparsity=None
137+
)
138+
sorting_analyzer.compute(extensions)
139+
# remove run_info.json to mimic a previous version of spikeinterface
140+
for ext in extensions:
141+
(folder / "extensions" / ext / "run_info.json").unlink()
142+
# should raise a warning for missing run_info
143+
with pytest.warns(UserWarning):
144+
sorting_analyzer = load_sorting_analyzer(folder, format="auto")
145+
146+
# zarr
147+
folder = tmp_path / "test_SortingAnalyzer_run_info.zarr"
148+
sorting_analyzer = create_sorting_analyzer(
149+
sorting, recording, format="zarr", folder=folder, sparse=False, sparsity=None
150+
)
151+
sorting_analyzer.compute(extensions)
152+
# remove run_info from attrs to mimic a previous version of spikeinterface
153+
root = sorting_analyzer._get_zarr_root(mode="r+")
154+
for ext in extensions:
155+
del root["extensions"][ext].attrs["run_info"]
156+
# should raise a warning for missing run_info
157+
with pytest.warns(UserWarning):
158+
sorting_analyzer = load_sorting_analyzer(folder, format="auto")
159+
160+
128161
def test_SortingAnalyzer_tmp_recording(dataset):
129162
recording, sorting = dataset
130163
recording_cached = recording.save(mode="memory")

0 commit comments

Comments
 (0)