xarray-behave 0.35.0__tar.gz → 0.35.2__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (57) hide show
  1. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/.github/workflows/publish.yaml +4 -3
  2. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/PKG-INFO +1 -1
  3. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/conda/xarray-behave/meta.yaml +2 -1
  4. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/__init__.py +1 -1
  5. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/annot.py +12 -19
  6. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/event_utils.py +7 -3
  7. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/gui/app.py +45 -134
  8. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/gui/das.py +8 -5
  9. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/gui/views.py +1 -1
  10. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/.gitignore +0 -0
  11. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/LICENSE +0 -0
  12. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/README.md +0 -0
  13. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/build_env.yml +0 -0
  14. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/conda/xarray-behave/bld.bat +0 -0
  15. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/conda/xarray-behave/build.sh +0 -0
  16. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/conda/xarray-behave/conda_build_config.yaml +0 -0
  17. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/condarc.yml +0 -0
  18. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/doc/demo.ipynb +0 -0
  19. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/doc/demo_behavioral_features.ipynb +0 -0
  20. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/doc/demo_behavioral_features_large_group.ipynb +0 -0
  21. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/doc/ncb.mplstyle +0 -0
  22. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/pyproject.toml +0 -0
  23. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/setup.py +0 -0
  24. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/gui/__init__.py +0 -0
  25. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/gui/audio_player.py +0 -0
  26. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/gui/formbuilder.py +0 -0
  27. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/gui/forms/das_make.yaml +0 -0
  28. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/gui/forms/das_predict.yaml +0 -0
  29. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/gui/forms/das_train.yaml +0 -0
  30. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/gui/forms/envelope_computation.yaml +0 -0
  31. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/gui/forms/export_for_das.yaml +0 -0
  32. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/gui/forms/from_dir.yaml +0 -0
  33. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/gui/forms/from_file.yaml +0 -0
  34. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/gui/forms/from_zarr.yaml +0 -0
  35. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/gui/forms/spec_freq.yaml +0 -0
  36. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/gui/icon.png +0 -0
  37. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/gui/table.py +0 -0
  38. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/gui/utils.py +0 -0
  39. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/gui/view_dialog.py +0 -0
  40. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/gui/widgets.py +0 -0
  41. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/io/__init__.py +0 -0
  42. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/io/annotations.py +0 -0
  43. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/io/annotations_manual.py +0 -0
  44. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/io/audio.py +0 -0
  45. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/io/balltracks.py +0 -0
  46. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/io/movieparams.py +0 -0
  47. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/io/poses.py +0 -0
  48. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/io/timestamps.py +0 -0
  49. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/io/tracks.py +0 -0
  50. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/loaders.py +0 -0
  51. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/metrics.py +0 -0
  52. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/src/xarray_behave/xarray_behave.py +0 -0
  53. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/tests/test_annot.py +0 -0
  54. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/tests/test_assemble.py +0 -0
  55. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/tests/test_assemble_metrics.py +0 -0
  56. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/tests/test_imports.py +0 -0
  57. {xarray-behave-0.35.0 → xarray-behave-0.35.2}/tests/test_io.py +0 -0
@@ -19,15 +19,16 @@ jobs:
19
19
  strategy:
20
20
  fail-fast: False
21
21
  matrix:
22
- python-version: [3.9 , '3.10', '3.11']
23
- os: [ubuntu-latest, windows-latest, macOS-13, macOS-14]
22
+ python-version: [3.9 , '3.10'] # , '3.11']
23
+ # os: [ubuntu-latest, windows-latest, macOS-13, macOS-14]
24
+ os: [macOS-13, macOS-14]
24
25
  # python-version: ['3.11']
25
26
  # os: [ubuntu-latest]
26
27
  defaults: # https://github.com/marketplace/actions/setup-miniconda#use-a-default-shell
27
28
  run:
28
29
  shell: bash -l {0}
29
30
  steps:
30
- - uses: actions/checkout@v3
31
+ - uses: actions/checkout@v4
31
32
  - name: Setup miniconda # https://github.com/marketplace/actions/setup-miniconda
32
33
  uses: conda-incubator/setup-miniconda@v3
33
34
  with:
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.1
2
2
  Name: xarray-behave
3
- Version: 0.35.0
3
+ Version: 0.35.2
4
4
  Summary: xarray tools for behavioral data.
5
5
  Author-email: Jan Clemens <clemensjan@googlemail.com>
6
6
  Requires-Python: >3.6
@@ -35,7 +35,8 @@ requirements:
35
35
  - xarray
36
36
  - dask
37
37
  #- py
38
- - conda-forge::pyside6 # [not win]
38
+ - conda-forge::pyside6 # [linux]
39
+ - conda-forge::pyside2 # [osx or arm64]
39
40
  - pyside6 # [win and py==310]
40
41
  - pyside2 # [win and py>310]
41
42
  - pyside2 # [win and py==39]
@@ -1,5 +1,5 @@
1
1
  """xarray tools for behavioral data."""
2
- __version__ = "0.35.0"
2
+ __version__ = "0.35.2"
3
3
 
4
4
  from .xarray_behave import assemble, assemble_metrics, load, save
5
5
  import os
@@ -132,19 +132,20 @@ class Events(UserDict):
132
132
  names = []
133
133
  start_seconds = []
134
134
  stop_seconds = []
135
-
136
135
  if len(dct.values()):
137
- if list(dct.values())[0].shape[1] > 2:
136
+ # check if there are annotations and if so, if they comes with channel information
137
+ if list(dct.values())[0].ndim > 1 and list(dct.values())[0].shape[1] > 2:
138
138
  channels = []
139
139
  else:
140
140
  channels = None
141
141
 
142
142
  for k, v in dct.items():
143
- names.extend([k] * v.shape[0])
144
- start_seconds.extend(v[:, 0])
145
- stop_seconds.extend(v[:, 1])
146
- if channels is not None:
147
- channels.extend(v[:, 2])
143
+ if len(v): # check if there are annotations
144
+ names.extend([k] * v.shape[0])
145
+ start_seconds.extend(v[:, 0])
146
+ stop_seconds.extend(v[:, 1])
147
+ if channels is not None:
148
+ channels.extend(v[:, 2])
148
149
  out = cls.from_lists(names, start_seconds, stop_seconds, channels=channels)
149
150
  return out
150
151
 
@@ -165,9 +166,7 @@ class Events(UserDict):
165
166
  columns.append("channel")
166
167
  return pd.DataFrame(columns=columns)
167
168
 
168
- def _append_row(
169
- self, df: pd.DataFrame, name: str, start_seconds: float, stop_seconds: Optional[float] = None, channel: int = -1
170
- ):
169
+ def _append_row(self, df: pd.DataFrame, name: str, start_seconds: float, stop_seconds: Optional[float] = None, channel: int = -1):
171
170
  if stop_seconds is None:
172
171
  stop_seconds = start_seconds
173
172
 
@@ -196,16 +195,12 @@ class Events(UserDict):
196
195
  """
197
196
  df = self._init_df()
198
197
  for name in self.names:
199
- for start_second, stop_second, channel in zip(
200
- self.start_seconds(name), self.stop_seconds(name), self.channels(name)
201
- ):
198
+ for start_second, stop_second, channel in zip(self.start_seconds(name), self.stop_seconds(name), self.channels(name)):
202
199
  df = self._append_row(df, name, start_second, stop_second, channel)
203
200
  if preserve_empty: # ensure we keep events without annotations
204
201
  for name, cat in zip(self.names, self.categories.values()):
205
202
  if name not in df.name.values:
206
- stop_seconds = (
207
- np.nan if cat == "event" else 0
208
- ) # (np.nan, np.nan) -> empty events, (np.nan, some number) -> empty segments
203
+ stop_seconds = np.nan if cat == "event" else 0 # (np.nan, np.nan) -> empty events, (np.nan, some number) -> empty segments
209
204
  df = self._append_row(df, name, start_seconds=np.nan, stop_seconds=stop_seconds)
210
205
  # make sure start and stop seconds are numeric
211
206
  df["start_seconds"] = pd.to_numeric(df["start_seconds"], errors="coerce")
@@ -315,9 +310,7 @@ class Events(UserDict):
315
310
  name = None
316
311
  return name
317
312
 
318
- def _get_index_of_nearest(
319
- self, time: float, name: str, tol: float = 0, min_time: Optional[float] = None, max_time: Optional[float] = None
320
- ):
313
+ def _get_index_of_nearest(self, time: float, name: str, tol: float = 0, min_time: Optional[float] = None, max_time: Optional[float] = None):
321
314
  within_range_indices = self.select_range(name, min_time, max_time, strict=False)
322
315
  if len(within_range_indices):
323
316
  nearest_start = self._find_nearest(self.start_seconds(name)[within_range_indices], time)
@@ -2,6 +2,7 @@ import logging
2
2
  import numpy as np
3
3
  import xarray as xr
4
4
  import pandas as pd
5
+ import scipy
5
6
 
6
7
 
7
8
  logger = logging.getLogger(__name__)
@@ -182,14 +183,14 @@ def eventtimes_to_traces(ds, event_times):
182
183
  return ds
183
184
 
184
185
 
185
- def traces_to_eventtimes(traces, event_names, event_categories):
186
+ def traces_to_eventtimes(traces, event_names, event_categories, events_are_binary: bool = True):
186
187
  """[summary]
187
188
 
188
189
  Args:
189
190
  traces ([type]): list of numpy arrays with the binary traces for each event/segment
190
191
  event_names ([type]): [description]
191
192
  event_categories ([type]): [description]
192
-
193
+ events_are_binary (bool, True): detect events indices where value is 1.0. Otherwise use scipy.signal.find_peaks.
193
194
  Returns:
194
195
  [type]: [description]
195
196
  """
@@ -202,7 +203,10 @@ def traces_to_eventtimes(traces, event_names, event_categories):
202
203
  for event_idx, (event_name, event_category) in enumerate(zip(event_names, event_categories)):
203
204
  logger.info(f" {event_name}")
204
205
  if event_category == "event":
205
- event_times[event_name] = np.where(traces[event_idx] == 1)[0]
206
+ if events_are_binary:
207
+ event_times[event_name] = np.where(traces[event_idx] == 1)[0]
208
+ else:
209
+ event_times[event_name], _ = scipy.signal.find_peaks(traces[event_idx])
206
210
  elif event_category == "segment":
207
211
  tmp = (traces[event_idx] == 1).astype(float) # makes this more robust
208
212
  onsets = np.where(np.diff(tmp) == 1)[0]
@@ -44,11 +44,7 @@ except ImportError:
44
44
  try:
45
45
  from . import das
46
46
  except Exception:
47
- logger.warning(
48
- "Failed to import the das module.\nIgnore if you do not want to use das.\n"
49
- "Otherwise follow these instructions to install:\n"
50
- "https://janclemenslab.org/das/install.html"
51
- )
47
+ logger.warning("Failed to import the das module.\nIgnore if you do not want to use das.\n" "Otherwise follow these instructions to install:\n" "https://janclemenslab.org/das/install.html")
52
48
 
53
49
  sys.setrecursionlimit(10**6) # increase recursion limit to avoid errors when keeping key pressed for a long time
54
50
  package_dir: str = xarray_behave.__path__[0]
@@ -157,9 +153,7 @@ class MainWindow(QtWidgets.QMainWindow):
157
153
 
158
154
  def save_swaps(self, qt_keycode=None):
159
155
  savefilename = self._get_filename_from_ds(suffix="_idswaps.txt")
160
- savefilename, _ = QtWidgets.QFileDialog.getSaveFileName(
161
- self, "Save swaps to", str(savefilename), filter="txt files (*.txt);;all files (*)"
162
- )
156
+ savefilename, _ = QtWidgets.QFileDialog.getSaveFileName(self, "Save swaps to", str(savefilename), filter="txt files (*.txt);;all files (*)")
163
157
  if len(savefilename):
164
158
  logger.info(f" Saving list of swap indices to {savefilename}.")
165
159
  os.makedirs(os.path.dirname(savefilename), exist_ok=True)
@@ -168,9 +162,7 @@ class MainWindow(QtWidgets.QMainWindow):
168
162
 
169
163
  def save_definitions(self, qt_keycode=None):
170
164
  savefilename = self._get_filename_from_ds(suffix="_definitions.csv")
171
- savefilename, _ = QtWidgets.QFileDialog.getSaveFileName(
172
- self, caption="Save definitions to", dir=str(savefilename), filter="CSV files (*_definitions.csv);;all files (*)"
173
- )
165
+ savefilename, _ = QtWidgets.QFileDialog.getSaveFileName(self, caption="Save definitions to", dir=str(savefilename), filter="CSV files (*_definitions.csv);;all files (*)")
174
166
  if len(savefilename):
175
167
  # get defs from annot and save them to csv
176
168
  logger.info(f" Saving definitions to {savefilename}.")
@@ -410,9 +402,7 @@ class MainWindow(QtWidgets.QMainWindow):
410
402
  self.export_to_h5(savefilename_trunk + ".h5", start_seconds, end_seconds) # , form_data["scale_audio"])
411
403
 
412
404
  logger.info(f" annotations to CSV: {savefilename_trunk + '.csv'}.")
413
- self.export_to_csv(
414
- savefilename_trunk + "_annotations.csv", start_seconds, end_seconds, which_events, match_to_samples=True
415
- )
405
+ self.export_to_csv(savefilename_trunk + "_annotations.csv", start_seconds, end_seconds, which_events, match_to_samples=True)
416
406
  logger.info("Done.")
417
407
 
418
408
  def das_make(self, qt_keycode=None):
@@ -513,9 +503,7 @@ class MainWindow(QtWidgets.QMainWindow):
513
503
  return form_data
514
504
 
515
505
  def save(arg):
516
- savefilename, _ = QtWidgets.QFileDialog.getSaveFileName(
517
- self, "Save configuration to", "", filter="yaml files (*.yaml);;all files (*)"
518
- )
506
+ savefilename, _ = QtWidgets.QFileDialog.getSaveFileName(self, "Save configuration to", "", filter="yaml files (*.yaml);;all files (*)")
519
507
  if len(savefilename):
520
508
  data = dialog.form.get_form_data()
521
509
  logger.info(f" Saving form fields to {savefilename}.")
@@ -526,9 +514,7 @@ class MainWindow(QtWidgets.QMainWindow):
526
514
  def make_cli(arg):
527
515
  script_ext = "cmd" if os.name == "nt" else "sh"
528
516
  savefilename = "train." + script_ext
529
- savefilename, _ = QtWidgets.QFileDialog.getSaveFileName(
530
- self, "Script name", savefilename, filter=f"script (*.{script_ext};;all files (*)"
531
- )
517
+ savefilename, _ = QtWidgets.QFileDialog.getSaveFileName(self, "Script name", savefilename, filter=f"script (*.{script_ext};;all files (*)")
532
518
  if len(savefilename):
533
519
  form_data = dialog.form.get_form_data()
534
520
  form_data = _filter_form_data(form_data, is_cli=True)
@@ -710,19 +696,13 @@ class MainWindow(QtWidgets.QMainWindow):
710
696
 
711
697
  params = das.utils.load_params(model_path)
712
698
  if audio.shape[0] < params["nb_hist"]:
713
- logger.warning(
714
- f" Aborting. Audio has fewer samples ({audio.shape[0]}) shorter"
715
- f" than network chunk size ({params['nb_hist']})."
716
- " Fix by select longer audio."
717
- )
699
+ logger.warning(f" Aborting. Audio has fewer samples ({audio.shape[0]}) shorter" f" than network chunk size ({params['nb_hist']})." " Fix by select longer audio.")
718
700
  return
719
701
 
720
702
  # select batch size so that at least 10 batches are run
721
703
  # minimizes loss of annotations from batch size "quantization" errors
722
704
  batch_size = 32
723
- nb_batches = lambda batch_size: int(
724
- np.floor((audio.shape[0] - ((batch_size - 1) + params["nb_hist"])) / (params["stride"] * (batch_size)))
725
- )
705
+ nb_batches = lambda batch_size: int(np.floor((audio.shape[0] - ((batch_size - 1) + params["nb_hist"])) / (params["stride"] * (batch_size))))
726
706
  while nb_batches(batch_size) < 10 and batch_size > 1:
727
707
  batch_size -= 1
728
708
 
@@ -785,11 +765,7 @@ class MainWindow(QtWidgets.QMainWindow):
785
765
  # segments['sequence'] = [s for s in segments['sequence'] if s is not None]
786
766
  detected_segment_names = np.unique(segments["sequence"])
787
767
  # if these are indices, get corresponding names
788
- if (
789
- len(detected_segment_names)
790
- and type(detected_segment_names[0]) is not str
791
- and type(detected_segment_names[0]) is not np.str_
792
- ):
768
+ if len(detected_segment_names) and type(detected_segment_names[0]) is not str and type(detected_segment_names[0]) is not np.str_:
793
769
  detected_segment_names = [segments["names"][ii] for ii in detected_segment_names]
794
770
 
795
771
  if len(detected_segment_names) > 0: # and detected_segment_names[0] is not None:
@@ -805,9 +781,7 @@ class MainWindow(QtWidgets.QMainWindow):
805
781
 
806
782
  onsets_seconds = self.ds.sampletime[onsets_samples]
807
783
  offsets_seconds = self.ds.sampletime[offsets_samples]
808
- for name_or_index, onset_seconds, offset_seconds in zip(
809
- segments["sequence"], onsets_seconds, offsets_seconds
810
- ):
784
+ for name_or_index, onset_seconds, offset_seconds in zip(segments["sequence"], onsets_seconds, offsets_seconds):
811
785
  if type(name_or_index) is not str and type(detected_segment_names[0]) is not np.str_:
812
786
  segment_name = segments["names"][name_or_index]
813
787
  else:
@@ -856,16 +830,8 @@ class MainWindow(QtWidgets.QMainWindow):
856
830
  except KeyError:
857
831
  pass
858
832
  except KeyError:
859
- logger.info(
860
- f"{filename} no sample rate info in NPZ file."
861
- f"Need to save 'samplerate' variable with the audio data. Defaulting to {samplerate}"
862
- )
863
- elif (
864
- filename.endswith(".h5")
865
- or filename.endswith(".hdfs")
866
- or filename.endswith(".hdf5")
867
- or filename.endswith(".mat")
868
- ):
833
+ logger.info(f"{filename} no sample rate info in NPZ file." f"Need to save 'samplerate' variable with the audio data. Defaulting to {samplerate}")
834
+ elif filename.endswith(".h5") or filename.endswith(".hdfs") or filename.endswith(".hdf5") or filename.endswith(".mat"):
869
835
  # infer data set (for hdf5) and populate form
870
836
  try:
871
837
  # list all data sets in file and add to list
@@ -969,9 +935,7 @@ class MainWindow(QtWidgets.QMainWindow):
969
935
  if not dirname:
970
936
  dirname = QtWidgets.QFileDialog.getExistingDirectory(parent=None, caption="Select data directory")
971
937
  if dirname:
972
- dialog = YamlDialog(
973
- yaml_file=package_dir + "/gui/forms/from_dir.yaml", title=f"Dataset from data directory {dirname}"
974
- )
938
+ dialog = YamlDialog(yaml_file=package_dir + "/gui/forms/from_dir.yaml", title=f"Dataset from data directory {dirname}")
975
939
 
976
940
  # initialize form data with cli args
977
941
  dialog.form["pixel_size_mm"] = pixel_size_mm # and un-disable
@@ -1046,9 +1010,7 @@ class MainWindow(QtWidgets.QMainWindow):
1046
1010
 
1047
1011
  # add event categories if they are missing in the dataset
1048
1012
  if "song_events" in ds and "event_categories" not in ds:
1049
- event_categories = [
1050
- "segment" if "sine" in evt or "syllable" in evt else "event" for evt in ds.event_types.values
1051
- ]
1013
+ event_categories = ["segment" if "sine" in evt or "syllable" in evt else "event" for evt in ds.event_types.values]
1052
1014
  ds = ds.assign_coords({"event_categories": (("event_types"), event_categories)})
1053
1015
 
1054
1016
  # add missing song types
@@ -1106,9 +1068,7 @@ class MainWindow(QtWidgets.QMainWindow):
1106
1068
  if not filename:
1107
1069
  filename, _ = QtWidgets.QFileDialog.getOpenFileName(parent=None, caption="Select dataset")
1108
1070
  if filename:
1109
- dialog = YamlDialog(
1110
- yaml_file=package_dir + "/gui/forms/from_zarr.yaml", title=f"Load dataset from zarr file {filename}"
1111
- )
1071
+ dialog = YamlDialog(yaml_file=package_dir + "/gui/forms/from_zarr.yaml", title=f"Load dataset from zarr file {filename}")
1112
1072
 
1113
1073
  # initialize form data with cli args
1114
1074
  if spec_freq_min is not None:
@@ -1146,9 +1106,7 @@ class MainWindow(QtWidgets.QMainWindow):
1146
1106
 
1147
1107
  # add event categories if they are missing in the dataset
1148
1108
  if "song_events" in ds and "event_categories" not in ds:
1149
- event_categories = [
1150
- "segment" if "sine" in evt or "syllable" in evt else "event" for evt in ds.event_types.values
1151
- ]
1109
+ event_categories = ["segment" if "sine" in evt or "syllable" in evt else "event" for evt in ds.event_types.values]
1152
1110
  ds = ds.assign_coords({"event_categories": (("event_types"), event_categories)})
1153
1111
  logger.info(ds)
1154
1112
  vr = None
@@ -1193,15 +1151,11 @@ class MainWindow(QtWidgets.QMainWindow):
1193
1151
 
1194
1152
  def save_dataset(self, qt_keycode=None):
1195
1153
  try:
1196
- savefilename = Path(
1197
- self.ds.attrs["root"], self.ds.attrs["dat_path"], self.ds.attrs["datename"], f"{self.ds.attrs['datename']}.zarr"
1198
- )
1154
+ savefilename = Path(self.ds.attrs["root"], self.ds.attrs["dat_path"], self.ds.attrs["datename"], f"{self.ds.attrs['datename']}.zarr")
1199
1155
  except KeyError:
1200
1156
  savefilename = ""
1201
1157
 
1202
- savefilename, _ = QtWidgets.QFileDialog.getSaveFileName(
1203
- self, "Save dataset to", str(savefilename), filter="zarr files (*.zarr);;all files (*)"
1204
- )
1158
+ savefilename, _ = QtWidgets.QFileDialog.getSaveFileName(self, "Save dataset to", str(savefilename), filter="zarr files (*.zarr);;all files (*)")
1205
1159
 
1206
1160
  if len(savefilename):
1207
1161
  file_exists = os.path.exists(savefilename)
@@ -1453,16 +1407,10 @@ class PSV(MainWindow):
1453
1407
  self._add_keyed_menuitem(view_video, "Change other fly", self.change_other_fly, "Z")
1454
1408
  self._add_keyed_menuitem(view_video, "Swap flies", self.swap_flies, "X")
1455
1409
  view_video.addSeparator()
1456
- self._add_keyed_menuitem(
1457
- view_video, "Move poses", partial(self.toggle, "move_poses"), "B", checkable=True, checked=self.move_poses
1458
- )
1410
+ self._add_keyed_menuitem(view_video, "Move poses", partial(self.toggle, "move_poses"), "B", checkable=True, checked=self.move_poses)
1459
1411
  view_video.addSeparator()
1460
- self._add_keyed_menuitem(
1461
- view_video, "Show fly position", partial(self.toggle, "show_dot"), "O", checkable=True, checked=self.show_dot
1462
- )
1463
- self._add_keyed_menuitem(
1464
- view_video, "Show poses", partial(self.toggle, "show_poses"), "P", checkable=True, checked=self.show_poses
1465
- )
1412
+ self._add_keyed_menuitem(view_video, "Show fly position", partial(self.toggle, "show_dot"), "O", checkable=True, checked=self.show_dot)
1413
+ self._add_keyed_menuitem(view_video, "Show poses", partial(self.toggle, "show_poses"), "P", checkable=True, checked=self.show_poses)
1466
1414
 
1467
1415
  view_audio = self.bar.addMenu("Audio")
1468
1416
  self._add_keyed_menuitem(view_audio, "Play waveform through speakers", self.play_audio, "E")
@@ -1486,9 +1434,7 @@ class PSV(MainWindow):
1486
1434
  self._add_keyed_menuitem(view_audio, "Select previous channel", self.set_next_channel, "Up")
1487
1435
  self._add_keyed_menuitem(view_audio, "Select next channel", self.set_prev_channel, "Down")
1488
1436
  view_audio.addSeparator()
1489
- self._add_keyed_menuitem(
1490
- view_audio, "Show spectrogram", partial(self.toggle, "show_spec"), None, checkable=True, checked=self.show_spec
1491
- )
1437
+ self._add_keyed_menuitem(view_audio, "Show spectrogram", partial(self.toggle, "show_spec"), None, checkable=True, checked=self.show_spec)
1492
1438
  self._add_keyed_menuitem(view_audio, "Increase frequency resolution", self.inc_freq_res, "R")
1493
1439
  self._add_keyed_menuitem(view_audio, "Increase temporal resolution", self.dec_freq_res, "T")
1494
1440
  view_audio.addSeparator()
@@ -1542,34 +1488,20 @@ class PSV(MainWindow):
1542
1488
  self._add_keyed_menuitem(view_annotations, "Generate proposal by envelope thresholding", self.threshold, "I")
1543
1489
  self._add_keyed_menuitem(view_annotations, "Adjust thresholding mode", self.set_envelope_computation)
1544
1490
  view_annotations.addSeparator()
1545
- self._add_keyed_menuitem(
1546
- view_annotations, "Approve proposals for active song type in view", self.approve_active_proposals, "G"
1547
- )
1548
- self._add_keyed_menuitem(
1549
- view_annotations, "Approve proposals for all song types in view", self.approve_all_proposals, "H"
1550
- )
1491
+ self._add_keyed_menuitem(view_annotations, "Approve proposals for active song type in view", self.approve_active_proposals, "G")
1492
+ self._add_keyed_menuitem(view_annotations, "Approve proposals for all song types in view", self.approve_all_proposals, "H")
1551
1493
 
1552
1494
  view_view = self.bar.addMenu("View")
1553
1495
  self._add_keyed_menuitem(view_view, "Video, waveform, and spectrogram display parameters", self.set_spec_freq)
1554
1496
  view_view.addSeparator()
1555
1497
  # TODO? only show these if tracks and/or video
1556
- self._add_keyed_menuitem(
1557
- view_view, "Show spectrogram", partial(self.toggle, "show_spec"), None, checkable=True, checked=self.show_spec
1558
- )
1559
- self._add_keyed_menuitem(
1560
- view_view, "Show waveform", partial(self.toggle, "show_trace"), None, checkable=True, checked=self.show_trace
1561
- )
1562
- self._add_keyed_menuitem(
1563
- view_view, "Show ethogram", partial(self.toggle, "show_annot"), None, checkable=True, checked=self.show_annot
1564
- )
1498
+ self._add_keyed_menuitem(view_view, "Show spectrogram", partial(self.toggle, "show_spec"), None, checkable=True, checked=self.show_spec)
1499
+ self._add_keyed_menuitem(view_view, "Show waveform", partial(self.toggle, "show_trace"), None, checkable=True, checked=self.show_trace)
1500
+ self._add_keyed_menuitem(view_view, "Show ethogram", partial(self.toggle, "show_annot"), None, checkable=True, checked=self.show_annot)
1565
1501
  if "pose_positions_allo" in self.ds:
1566
- self._add_keyed_menuitem(
1567
- view_view, "Show tracks", partial(self.toggle, "show_tracks"), None, checkable=True, checked=self.show_tracks
1568
- )
1502
+ self._add_keyed_menuitem(view_view, "Show tracks", partial(self.toggle, "show_tracks"), None, checkable=True, checked=self.show_tracks)
1569
1503
  if self.vr is not None:
1570
- self._add_keyed_menuitem(
1571
- view_view, "Show movie", partial(self.toggle, "show_movie"), None, checkable=True, checked=self.show_movie
1572
- )
1504
+ self._add_keyed_menuitem(view_view, "Show movie", partial(self.toggle, "show_movie"), None, checkable=True, checked=self.show_movie)
1573
1505
 
1574
1506
  self.hl = QtWidgets.QHBoxLayout()
1575
1507
 
@@ -1940,9 +1872,7 @@ class PSV(MainWindow):
1940
1872
 
1941
1873
  def delete_current_events(self, qt_keycode):
1942
1874
  if self.current_event_index is not None:
1943
- deleted_events = self.event_times.delete_range(
1944
- self.current_event_name, self.time0 / self.fs_song, self.time1 / self.fs_song
1945
- )
1875
+ deleted_events = self.event_times.delete_range(self.current_event_name, self.time0 / self.fs_song, self.time1 / self.fs_song)
1946
1876
  nb_deleted_events = len(deleted_events)
1947
1877
  if nb_deleted_events:
1948
1878
  logger.info(f" Deleted {nb_deleted_events} annotation(s) of type {self.current_event_name}.")
@@ -1964,9 +1894,7 @@ class PSV(MainWindow):
1964
1894
  def threshold(self, qt_keycode):
1965
1895
  if self.STOP and self.current_event_name is not None:
1966
1896
  if self.event_times.categories[self.current_event_name] == "event":
1967
- indexes = peakutils.indexes(
1968
- self.envelope, thres=self.slice_view.threshold, min_dist=self.thres_min_dist * self.fs_song, thres_abs=True
1969
- )
1897
+ indexes = peakutils.indexes(self.envelope, thres=self.slice_view.threshold, min_dist=self.thres_min_dist * self.fs_song, thres_abs=True)
1970
1898
  # add events to current song type
1971
1899
  for t in self.x[indexes]:
1972
1900
  self.event_times.add_time(self.current_event_name, t)
@@ -2116,9 +2044,7 @@ class PSV(MainWindow):
2116
2044
  dialog.exec_()
2117
2045
 
2118
2046
  def set_envelope_computation(self, qt_keycode):
2119
- dialog = YamlDialog(
2120
- yaml_file=package_dir + "/gui/forms/envelope_computation.yaml", title="Set options for envelope computation"
2121
- )
2047
+ dialog = YamlDialog(yaml_file=package_dir + "/gui/forms/envelope_computation.yaml", title="Set options for envelope computation")
2122
2048
 
2123
2049
  dialog.form["thres_min_dist"] = self.thres_min_dist
2124
2050
  dialog.form["thres_env_std"] = self.thres_env_std
@@ -2196,9 +2122,7 @@ class PSV(MainWindow):
2196
2122
  i1 = int(self.time1 / self.fs_ratio)
2197
2123
 
2198
2124
  self.x_tracks = self.ds.time.data[i0:i1]
2199
- self.y_tracks = self.ds.pose_positions_allo.data[
2200
- i0:i1, self.focal_fly, self.track_sel_names, self.track_sel_coords
2201
- ]
2125
+ self.y_tracks = self.ds.pose_positions_allo.data[i0:i1, self.focal_fly, self.track_sel_names, self.track_sel_coords]
2202
2126
  self.tracks_view.update_trace()
2203
2127
  self.tracks_view.show()
2204
2128
  else:
@@ -2245,13 +2169,9 @@ class PSV(MainWindow):
2245
2169
  if self.event_times.categories[event_name] == "segment":
2246
2170
  for onset, offset in zip(events_in_view[:, 0], events_in_view[:, 1]):
2247
2171
  if self.show_trace:
2248
- self.slice_view.add_segment(
2249
- onset, offset, event_index, brush=event_brush, pen=event_pen, movable=movable, text=segment_text
2250
- )
2172
+ self.slice_view.add_segment(onset, offset, event_index, brush=event_brush, pen=event_pen, movable=movable, text=segment_text)
2251
2173
  if self.show_tracks:
2252
- self.tracks_view.add_segment(
2253
- onset, offset, event_index, brush=event_brush, pen=event_pen, movable=movable, text=segment_text
2254
- )
2174
+ self.tracks_view.add_segment(onset, offset, event_index, brush=event_brush, pen=event_pen, movable=movable, text=segment_text)
2255
2175
  if self.show_annot:
2256
2176
  self.annot_view.add_segment(
2257
2177
  onset,
@@ -2263,9 +2183,7 @@ class PSV(MainWindow):
2263
2183
  text=segment_text,
2264
2184
  )
2265
2185
  if self.show_spec:
2266
- self.spec_view.add_segment(
2267
- onset, offset, event_index, brush=event_brush, pen=event_pen, movable=movable, text=segment_text
2268
- )
2186
+ self.spec_view.add_segment(onset, offset, event_index, brush=event_brush, pen=event_pen, movable=movable, text=segment_text)
2269
2187
  elif self.event_times.categories[event_name] == "event":
2270
2188
  if self.show_trace:
2271
2189
  self.slice_view.add_event(events_in_view[:, 0], event_index, event_pen, movable=movable, text=segment_text)
@@ -2298,9 +2216,7 @@ class PSV(MainWindow):
2298
2216
 
2299
2217
  new_region = region.getRegion()
2300
2218
  self.event_times.move_time(event_name_to_move, region.bounds, new_region)
2301
- logger.info(
2302
- f" Moved {event_name_to_move} from t=[{region.bounds[0]:1.4f}:{region.bounds[1]:1.4f}] to [{new_region[0]:1.4f}:{new_region[1]:1.4f}] seconds."
2303
- )
2219
+ logger.info(f" Moved {event_name_to_move} from t=[{region.bounds[0]:1.4f}:{region.bounds[1]:1.4f}] to [{new_region[0]:1.4f}:{new_region[1]:1.4f}] seconds.")
2304
2220
 
2305
2221
  mp = self.annot_view.mousePoint.y()
2306
2222
  if mp > 0 and mp < 1:
@@ -2369,10 +2285,7 @@ class PSV(MainWindow):
2369
2285
  fly_pos = self.ds.pose_positions_allo.data[self.index_other, :, self.pose_center_index, :]
2370
2286
  fly_pos = np.array(fly_pos) # in case this is a dask.array
2371
2287
  if self.crop: # transform fly pos to coordinates of the cropped box
2372
- box_center = (
2373
- self.ds.pose_positions_allo.data[self.index_other, self.focal_fly, self.pose_center_index]
2374
- + self.box_size / 2
2375
- )
2288
+ box_center = self.ds.pose_positions_allo.data[self.index_other, self.focal_fly, self.pose_center_index] + self.box_size / 2
2376
2289
  box_center = np.array(box_center) # in case this is a dask.array
2377
2290
  fly_pos = fly_pos - box_center
2378
2291
  fly_dist = np.sum((fly_pos - np.array([mouseY, mouseX])) ** 2, axis=-1)
@@ -2412,9 +2325,7 @@ class PSV(MainWindow):
2412
2325
  if self.event_times.categories[self.current_event_name] == "event":
2413
2326
  logger.info(f" Changed event at {changed_time[0]:1.4f} from {old_name} to {new_name}.")
2414
2327
  else:
2415
- logger.info(
2416
- f" Changed segment at {changed_time[0]:1.4f}:{changed_time[1]:1.4f} from {old_name} to {new_name}."
2417
- )
2328
+ logger.info(f" Changed segment at {changed_time[0]:1.4f}:{changed_time[1]:1.4f} from {old_name} to {new_name}.")
2418
2329
  self.update_xy()
2419
2330
  elif mouseButton == 1: # add event
2420
2331
  if self.current_event_index is not None:
@@ -2434,16 +2345,12 @@ class PSV(MainWindow):
2434
2345
  stop_seconds=mouseT,
2435
2346
  channel=self.current_channel_index,
2436
2347
  )
2437
- logger.info(
2438
- f" Added {self.current_event_name} on channel {self.current_channel_index} at t=[{self.sinet0:1.4f}:{mouseT:1.4f}] seconds."
2439
- )
2348
+ logger.info(f" Added {self.current_event_name} on channel {self.current_channel_index} at t=[{self.sinet0:1.4f}:{mouseT:1.4f}] seconds.")
2440
2349
  self.sinet0 = None
2441
2350
  if self.event_times.categories[self.current_event_name] == "event":
2442
2351
  self.sinet0 = None
2443
2352
  self.event_times.add_time(self.current_event_name, start_seconds=mouseT, channel=self.current_channel_index)
2444
- logger.info(
2445
- f" Added {self.current_event_name} on channel {self.current_channel_index} at t={mouseT:1.4f} seconds."
2446
- )
2353
+ logger.info(f" Added {self.current_event_name} on channel {self.current_channel_index} at t={mouseT:1.4f} seconds.")
2447
2354
  self.update_xy()
2448
2355
  else:
2449
2356
  self.sinet0 = None
@@ -2807,3 +2714,7 @@ def cli():
2807
2714
  logger.getLogger().setLevel(logging.INFO)
2808
2715
 
2809
2716
  defopt.run(main, show_defaults=False)
2717
+
2718
+
2719
+ if __name__ == "__main__":
2720
+ main_das()
@@ -8,6 +8,7 @@ import zarr
8
8
  import numpy as np
9
9
  import pandas as pd
10
10
  import librosa
11
+ import scipy
11
12
  import h5py
12
13
 
13
14
  import das.make_dataset as dsm
@@ -393,9 +394,7 @@ def make(
393
394
  store.attrs["data_splits"] = data_split_dict
394
395
  logger.info("Done.")
395
396
  # report
396
- logger.info(
397
- f" Got {store['train']['x'].shape}, {store['val']['x'].shape}, {store['test']['x'].shape} train/val/test samples."
398
- )
397
+ logger.info(f" Got {store['train']['x'].shape}, {store['val']['x'].shape}, {store['test']['x'].shape} train/val/test samples.")
399
398
  if to_npy_dir: # save as npy_dir
400
399
  logger.info(f" Saving to {store_folder}.")
401
400
  das.npy_dir.save(store_folder, store)
@@ -408,8 +407,12 @@ def make(
408
407
  names = store.attrs["class_names"][1:] # [1:] ignore the noise class
409
408
  types = store.attrs["class_types"][1:]
410
409
  for key in store.keys():
411
- ann = event_utils.traces_to_eventtimes(store[key]["y"][:, 1:].T, names, types)
412
- ann = {k: v / fs for k, v in ann.items()} # convert indices to seconds
410
+ ann = event_utils.traces_to_eventtimes(store[key]["y"][:, 1:].T, names, types, events_are_binary=False)
411
+ for k, v in ann.items():
412
+ if v.ndim == 1: # add start/stop seconds for events - traces_to_eventtimes only returns event start times but annot expects both start and end
413
+ v = np.stack((v, v)).T
414
+ v = v / fs # convert indices to seconds
415
+ ann[k] = v
413
416
  ev = annot.Events.from_dict(ann)
414
417
  ev.to_df().to_csv(os.path.join(store.attrs["store_folder"], key, "x_annotations.csv"))
415
418
  logger.info("The dataset has been made.")
@@ -584,7 +584,7 @@ class AnnotView(pg.PlotWidget):
584
584
  self.scene().sigMouseMoved.connect(self.mouseMoved)
585
585
 
586
586
  self.annotation_items = []
587
- self.getAxis("left").setStyle(tickFont=self.m.font_condensed)
587
+ # self.getAxis("left").setStyle(tickFont=self.m.font_condensed)
588
588
  self.getAxis("left").setWidth(50)
589
589
 
590
590
  self.mousePoint = None
File without changes
File without changes
File without changes