1111import shutil
1212import warnings
1313import importlib
14+ from time import perf_counter
1415
1516import 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
0 commit comments