Skip to content

Commit 70e70ef

Browse files
committed
transform: be lazier with the featurization in fit()
1 parent f9c518b commit 70e70ef

6 files changed

Lines changed: 20 additions & 14 deletions

File tree

src/dartsort/transform/amplitudes.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@ class AmplitudeFeatures(BaseWaveformFeaturizer):
1010
"""Extract spike amplitudes."""
1111

1212
is_multi = True
13+
featurize_in_fit = True
1314

1415
def __init__(
1516
self,
@@ -122,6 +123,7 @@ def transform(self, waveforms, **unused):
122123

123124
class AmplitudeVector(BaseWaveformFeaturizer):
124125
default_name = "amplitude_vectors"
126+
featurize_in_fit = True
125127

126128
def __init__(
127129
self,
@@ -153,6 +155,7 @@ def transform(self, waveforms, **unused):
153155

154156
class MaxAmplitude(BaseWaveformFeaturizer):
155157
default_name = "amplitudes"
158+
featurize_in_fit = True
156159
shape = ()
157160

158161
def __init__(

src/dartsort/transform/mixture_classifier.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -12,9 +12,10 @@
1212
"""
1313

1414
import sys
15+
from collections.abc import Sequence
1516
from dataclasses import replace
1617
from math import fabs
17-
from typing import TYPE_CHECKING, Sequence
18+
from typing import TYPE_CHECKING
1819

1920
import numpy as np
2021
import torch

src/dartsort/transform/pipeline.py

Lines changed: 6 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1,8 +1,8 @@
11
"""A class which manages pipelines of denoisers and featurizers"""
22

3+
from collections.abc import Sequence
34
from copy import deepcopy
45
from pathlib import Path
5-
from typing import Sequence
66

77
import torch
88
from 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 = {

src/dartsort/transform/reduction.py

Lines changed: 3 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -48,18 +48,15 @@ def __init__(
4848
pfx = ""
4949
else:
5050
pfx = name_prefix + "_"
51-
if not self.online:
52-
names = [f"{pfx}{_my_name}"]
53-
else:
54-
names = []
51+
names = [] if self.online else [f"{pfx}{_my_name}"]
5552
super().__init__(geom=geom, channel_index=channel_index, name=names)
5653
self.name_prefix = name_prefix
5754
self.reduction = reduction
5855
self.n_units = n_units
5956
self.feature_dim = feature_dim
6057
self.output_channels = output_channels
61-
self.shape = [(self.feature_dim, output_channels)]
62-
self.dtype = [dtype]
58+
self.shape = [] if self.online else [(self.feature_dim, output_channels)]
59+
self.dtype = [] if self.online else [dtype]
6360
self.with_raw_std_dev = with_raw_std_dev
6461
self._initialize((self.feature_dim, output_channels))
6562

src/dartsort/transform/transform_base.py

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,6 @@
1+
from collections.abc import Iterable
12
from pathlib import Path
2-
from typing import TYPE_CHECKING, Any, Iterable, Self
3+
from typing import TYPE_CHECKING, Any, Self
34

45
import torch
56
from spikeinterface.core import BaseRecording
@@ -25,6 +26,7 @@ class BaseWaveformModule(BModule):
2526
needs_residual = False
2627
fits_from_disk = False
2728
needs_more_features = False
29+
featurize_in_fit = False
2830

2931
def __init__(
3032
self,
@@ -231,7 +233,7 @@ def spike_datasets(self, force_save: bool = False) -> Iterable[SpikeDataset]:
231233
shape_per_spike=s,
232234
dtype=str(d).split(".")[1],
233235
)
234-
for n, s, d in zip(self.name, self.shape, self.dtype)
236+
for n, s, d in zip(self.name, self.shape, self.dtype, strict=True)
235237
]
236238
return datasets
237239
else:
@@ -321,6 +323,7 @@ def forward(self, waveforms, **spike_data):
321323

322324
class Waveform(BaseWaveformFeaturizer):
323325
default_name = "waveforms"
326+
featurize_in_fit = True
324327

325328
def __init__(
326329
self,

tests/test_clustering.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -86,6 +86,7 @@ def test_clustering(simulations, sim_name, featkw, cluskw):
8686
if cluskw["cluster_strategy"] == "density_peaks_uhdversion":
8787
if not featkw["use_amplitude"]:
8888
return
89+
featkw = featkw | {"need_xyza": True}
8990
sim = simulations[sim_name]
9091
recording = sim["recording"]
9192
sorting = sim["sorting"]

0 commit comments

Comments
 (0)