11"""A class which manages pipelines of denoisers and featurizers"""
22
3+ from collections .abc import Sequence
34from copy import deepcopy
45from pathlib import Path
5- from typing import Sequence
66
77import torch
88from spikeinterface .core import BaseRecording
@@ -341,7 +341,7 @@ def fit(
341341 if transformer .is_featurizer and transformer .is_denoiser :
342342 waveforms , new_features = transformer (** features )
343343 features .update (waveforms = waveforms , ** new_features )
344- elif transformer .is_featurizer :
344+ elif transformer .is_featurizer and transformer . featurize_in_fit :
345345 assert isinstance (transformer , BaseWaveformFeaturizer )
346346 features .update (transformer .transform (** features ))
347347 elif transformer .is_denoiser :
@@ -372,7 +372,7 @@ def precompute(self):
372372 def transform_to_disk (
373373 self ,
374374 hdf5_filename : str | Path ,
375- waveforms_dataset_name : str | None = "waveforms" ,
375+ waveforms_dataset_name : str = "waveforms" ,
376376 other_dset_names : Sequence [str ] | None = None ,
377377 start_index : int | None = None ,
378378 up_to_index : int | None = None ,
@@ -409,7 +409,6 @@ def transform_to_disk(
409409 with File (hdf5_filename , mode = "r+" , libver = "latest" , locking = False ) as h5 :
410410 if all (ds .name in h5 for ds in datasets ):
411411 return
412- wfs = h5 [waveforms_dataset_name ] if waveforms_dataset_name else None
413412 other_dsets = {od : h5 [od ] for od in other_dset_names }
414413 outs = {
415414 ds .name : h5 .create_dataset (
@@ -418,7 +417,9 @@ def transform_to_disk(
418417 for ds in datasets
419418 }
420419 for sli , chk in yield_chunks (
421- wfs , desc_prefix = "Transform to disk" , show_progress = False
420+ h5 [waveforms_dataset_name ],
421+ desc_prefix = "Transform to disk" ,
422+ show_progress = False ,
422423 ):
423424 chk_fp = {k : v [sli ].to (device = dev ) for k , v in fixed_properties .items ()}
424425 other_fp = {
0 commit comments