xarray-behave 0.35.2__tar.gz → 0.35.4__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.2 → xarray-behave-0.35.4}/.github/workflows/publish.yaml +5 -5
  2. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/PKG-INFO +1 -1
  3. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/conda/xarray-behave/meta.yaml +8 -10
  4. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/__init__.py +2 -1
  5. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/annot.py +27 -10
  6. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/gui/app.py +155 -58
  7. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/gui/das.py +6 -2
  8. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/gui/formbuilder.py +1 -0
  9. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/gui/views.py +9 -3
  10. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/io/audio.py +1 -0
  11. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/io/balltracks.py +1 -0
  12. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/io/movieparams.py +1 -0
  13. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/io/poses.py +7 -2
  14. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/loaders.py +1 -0
  15. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/metrics.py +1 -0
  16. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/xarray_behave.py +1 -0
  17. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/.gitignore +0 -0
  18. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/LICENSE +0 -0
  19. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/README.md +0 -0
  20. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/build_env.yml +0 -0
  21. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/conda/xarray-behave/bld.bat +0 -0
  22. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/conda/xarray-behave/build.sh +0 -0
  23. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/conda/xarray-behave/conda_build_config.yaml +0 -0
  24. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/condarc.yml +0 -0
  25. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/doc/demo.ipynb +0 -0
  26. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/doc/demo_behavioral_features.ipynb +0 -0
  27. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/doc/demo_behavioral_features_large_group.ipynb +0 -0
  28. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/doc/ncb.mplstyle +0 -0
  29. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/pyproject.toml +0 -0
  30. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/setup.py +0 -0
  31. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/event_utils.py +0 -0
  32. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/gui/__init__.py +0 -0
  33. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/gui/audio_player.py +0 -0
  34. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/gui/forms/das_make.yaml +0 -0
  35. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/gui/forms/das_predict.yaml +0 -0
  36. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/gui/forms/das_train.yaml +0 -0
  37. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/gui/forms/envelope_computation.yaml +0 -0
  38. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/gui/forms/export_for_das.yaml +0 -0
  39. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/gui/forms/from_dir.yaml +0 -0
  40. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/gui/forms/from_file.yaml +0 -0
  41. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/gui/forms/from_zarr.yaml +0 -0
  42. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/gui/forms/spec_freq.yaml +0 -0
  43. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/gui/icon.png +0 -0
  44. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/gui/table.py +0 -0
  45. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/gui/utils.py +0 -0
  46. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/gui/view_dialog.py +0 -0
  47. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/gui/widgets.py +0 -0
  48. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/io/__init__.py +0 -0
  49. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/io/annotations.py +0 -0
  50. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/io/annotations_manual.py +0 -0
  51. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/io/timestamps.py +0 -0
  52. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/src/xarray_behave/io/tracks.py +0 -0
  53. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/tests/test_annot.py +0 -0
  54. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/tests/test_assemble.py +0 -0
  55. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/tests/test_assemble_metrics.py +0 -0
  56. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/tests/test_imports.py +0 -0
  57. {xarray-behave-0.35.2 → xarray-behave-0.35.4}/tests/test_io.py +0 -0
@@ -19,11 +19,11 @@ 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]
24
- os: [macOS-13, macOS-14]
25
- # python-version: ['3.11']
26
- # os: [ubuntu-latest]
22
+ python-version: ['3.10']
23
+ os: [ubuntu-latest, windows-latest, macOS-13, macOS-14]
24
+ include:
25
+ - python-version: 3.9
26
+ os: windows-latest
27
27
  defaults: # https://github.com/marketplace/actions/setup-miniconda#use-a-default-shell
28
28
  run:
29
29
  shell: bash -l {0}
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.1
2
2
  Name: xarray-behave
3
- Version: 0.35.2
3
+ Version: 0.35.4
4
4
  Summary: xarray tools for behavioral data.
5
5
  Author-email: Jan Clemens <clemensjan@googlemail.com>
6
6
  Requires-Python: >3.6
@@ -16,6 +16,7 @@ requirements:
16
16
  host:
17
17
  - python {{ python }}
18
18
  - pip
19
+ - numpy=1.23
19
20
  run:
20
21
  - python {{ python }}
21
22
  - defopt=6.3
@@ -24,30 +25,27 @@ requirements:
24
25
  - h5py
25
26
  - librosa>0.8
26
27
  - matplotlib
27
- - matplotlib-scalebar
28
28
  - pandas
29
29
  - scipy>=1.9
30
30
  - peakutils
31
31
  - pyyaml
32
32
  - scikit-learn
33
33
  - zarr
34
- - numba # >=0.56
34
+ - numba
35
35
  - xarray
36
36
  - dask
37
- #- py
38
- - conda-forge::pyside6 # [linux]
39
- - conda-forge::pyside2 # [osx or arm64]
40
- - pyside6 # [win and py==310]
41
- - pyside2 # [win and py>310]
42
- - pyside2 # [win and py==39]
43
- - pyqtgraph>0.12.2
37
+ - conda-forge::pyside2
38
+ # - conda-forge::pyside2 # [linux or osx or arm64]
39
+ # - conda-forge::pyside6 # [win and py>39]
40
+ # - conda-forge::pyside2 # [osx or arm64]
41
+ # - pyside2 # [win and py==39]
42
+ - pyqtgraph>0.12
44
43
  - qtpy
45
44
  - superqt
46
45
  - rich
47
46
  - colorcet
48
47
  - python-sounddevice
49
48
  - scikit-image
50
- # - opencv
51
49
  - ffmpeg
52
50
  - pyvideoreader
53
51
  - samplestamps>=0.6
@@ -1,5 +1,6 @@
1
1
  """xarray tools for behavioral data."""
2
- __version__ = "0.35.2"
2
+
3
+ __version__ = "0.35.4"
3
4
 
4
5
  from .xarray_behave import assemble, assemble_metrics, load, save
5
6
  import os
@@ -2,6 +2,7 @@
2
2
 
3
3
  TODO: From/to traces
4
4
  """
5
+
5
6
  import numpy as np
6
7
  import xarray as xr
7
8
  import pandas as pd
@@ -166,7 +167,9 @@ class Events(UserDict):
166
167
  columns.append("channel")
167
168
  return pd.DataFrame(columns=columns)
168
169
 
169
- def _append_row(self, df: pd.DataFrame, name: str, start_seconds: float, stop_seconds: Optional[float] = None, channel: int = -1):
170
+ def _append_row(
171
+ self, df: pd.DataFrame, name: str, start_seconds: float, stop_seconds: Optional[float] = None, channel: int = -1
172
+ ):
170
173
  if stop_seconds is None:
171
174
  stop_seconds = start_seconds
172
175
 
@@ -195,12 +198,16 @@ class Events(UserDict):
195
198
  """
196
199
  df = self._init_df()
197
200
  for name in self.names:
198
- for start_second, stop_second, channel in zip(self.start_seconds(name), self.stop_seconds(name), self.channels(name)):
201
+ for start_second, stop_second, channel in zip(
202
+ self.start_seconds(name), self.stop_seconds(name), self.channels(name)
203
+ ):
199
204
  df = self._append_row(df, name, start_second, stop_second, channel)
200
205
  if preserve_empty: # ensure we keep events without annotations
201
206
  for name, cat in zip(self.names, self.categories.values()):
202
207
  if name not in df.name.values:
203
- stop_seconds = np.nan if cat == "event" else 0 # (np.nan, np.nan) -> empty events, (np.nan, some number) -> empty segments
208
+ stop_seconds = (
209
+ np.nan if cat == "event" else 0
210
+ ) # (np.nan, np.nan) -> empty events, (np.nan, some number) -> empty segments
204
211
  df = self._append_row(df, name, start_seconds=np.nan, stop_seconds=stop_seconds)
205
212
  # make sure start and stop seconds are numeric
206
213
  df["start_seconds"] = pd.to_numeric(df["start_seconds"], errors="coerce")
@@ -310,7 +317,9 @@ class Events(UserDict):
310
317
  name = None
311
318
  return name
312
319
 
313
- def _get_index_of_nearest(self, time: float, name: str, tol: float = 0, min_time: Optional[float] = None, max_time: Optional[float] = None):
320
+ def _get_index_of_nearest(
321
+ self, time: float, name: str, tol: float = 0, min_time: Optional[float] = None, max_time: Optional[float] = None
322
+ ):
314
323
  within_range_indices = self.select_range(name, min_time, max_time, strict=False)
315
324
  if len(within_range_indices):
316
325
  nearest_start = self._find_nearest(self.start_seconds(name)[within_range_indices], time)
@@ -338,9 +347,9 @@ class Events(UserDict):
338
347
  event_at_time = matching_start < time
339
348
  elif self.categories[name] == "event":
340
349
  if nearest_is_start:
341
- event_at_time = np.abs(time - nearest_start) < tol
350
+ event_at_time = np.abs(time - nearest_start) <= tol
342
351
  else:
343
- event_at_time = np.abs(time - nearest_stop) < tol
352
+ event_at_time = np.abs(time - nearest_stop) <= tol
344
353
  else:
345
354
  event_at_time = False
346
355
 
@@ -350,7 +359,13 @@ class Events(UserDict):
350
359
  return index
351
360
 
352
361
  def change_name(
353
- self, time: float, new_name: str, tol: float = 0, min_time: Optional[float] = None, max_time: Optional[float] = None
362
+ self,
363
+ time: float,
364
+ new_name: str,
365
+ tol: float = 0,
366
+ min_time: Optional[float] = None,
367
+ max_time: Optional[float] = None,
368
+ old_name: Optional[str] = None,
354
369
  ) -> Tuple[Optional[List[int]], Optional[str], Optional[str]]:
355
370
  """Change the name of the annotation.
356
371
 
@@ -360,19 +375,21 @@ class Events(UserDict):
360
375
  tol (float, optional): Tolerance for matching events. Defaults to 0.
361
376
  min_time (Optional[float], optional): _description_. Defaults to None.
362
377
  max_time (Optional[float], optional): _description_. Defaults to None.
363
-
378
+ old_name (Optional[str]): name of event to move. Defaults to None.
364
379
  Returns:
365
380
  Tuple[List[int], str, str]: ([start_seconds, stop_seconds], old_name, new_name
366
381
  Tuple[None, None, None] if no event near time, or new_name is old_name
367
382
  """
368
- name = self._get_name_of_nearest(time, min_time, max_time)
383
+ if old_name is None:
384
+ name = self._get_name_of_nearest(time, min_time, max_time)
385
+ else:
386
+ name = old_name
369
387
 
370
388
  # nothing to do
371
389
  if name is None or name == new_name or self.categories[name] != self.categories[new_name]:
372
390
  return None, None, None
373
391
 
374
392
  index = self._get_index_of_nearest(time, name, tol, min_time, max_time)
375
- print(index)
376
393
  if index is not None:
377
394
  changed_time = self[name][index, :]
378
395
  old_name = name
@@ -2,6 +2,7 @@
2
2
 
3
3
  `python -m xarray_behave.ui datename root`
4
4
  """
5
+
5
6
  import os
6
7
  import sys
7
8
  import logging
@@ -44,7 +45,11 @@ except ImportError:
44
45
  try:
45
46
  from . import das
46
47
  except Exception:
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")
48
+ logger.warning(
49
+ "Failed to import the das module.\nIgnore if you do not want to use das.\n"
50
+ "Otherwise follow these instructions to install:\n"
51
+ "https://janclemenslab.org/das/install.html"
52
+ )
48
53
 
49
54
  sys.setrecursionlimit(10**6) # increase recursion limit to avoid errors when keeping key pressed for a long time
50
55
  package_dir: str = xarray_behave.__path__[0]
@@ -153,7 +158,9 @@ class MainWindow(QtWidgets.QMainWindow):
153
158
 
154
159
  def save_swaps(self, qt_keycode=None):
155
160
  savefilename = self._get_filename_from_ds(suffix="_idswaps.txt")
156
- savefilename, _ = QtWidgets.QFileDialog.getSaveFileName(self, "Save swaps to", str(savefilename), filter="txt files (*.txt);;all files (*)")
161
+ savefilename, _ = QtWidgets.QFileDialog.getSaveFileName(
162
+ self, "Save swaps to", str(savefilename), filter="txt files (*.txt);;all files (*)"
163
+ )
157
164
  if len(savefilename):
158
165
  logger.info(f" Saving list of swap indices to {savefilename}.")
159
166
  os.makedirs(os.path.dirname(savefilename), exist_ok=True)
@@ -162,7 +169,9 @@ class MainWindow(QtWidgets.QMainWindow):
162
169
 
163
170
  def save_definitions(self, qt_keycode=None):
164
171
  savefilename = self._get_filename_from_ds(suffix="_definitions.csv")
165
- savefilename, _ = QtWidgets.QFileDialog.getSaveFileName(self, caption="Save definitions to", dir=str(savefilename), filter="CSV files (*_definitions.csv);;all files (*)")
172
+ savefilename, _ = QtWidgets.QFileDialog.getSaveFileName(
173
+ self, caption="Save definitions to", dir=str(savefilename), filter="CSV files (*_definitions.csv);;all files (*)"
174
+ )
166
175
  if len(savefilename):
167
176
  # get defs from annot and save them to csv
168
177
  logger.info(f" Saving definitions to {savefilename}.")
@@ -402,7 +411,9 @@ class MainWindow(QtWidgets.QMainWindow):
402
411
  self.export_to_h5(savefilename_trunk + ".h5", start_seconds, end_seconds) # , form_data["scale_audio"])
403
412
 
404
413
  logger.info(f" annotations to CSV: {savefilename_trunk + '.csv'}.")
405
- self.export_to_csv(savefilename_trunk + "_annotations.csv", start_seconds, end_seconds, which_events, match_to_samples=True)
414
+ self.export_to_csv(
415
+ savefilename_trunk + "_annotations.csv", start_seconds, end_seconds, which_events, match_to_samples=True
416
+ )
406
417
  logger.info("Done.")
407
418
 
408
419
  def das_make(self, qt_keycode=None):
@@ -503,7 +514,9 @@ class MainWindow(QtWidgets.QMainWindow):
503
514
  return form_data
504
515
 
505
516
  def save(arg):
506
- savefilename, _ = QtWidgets.QFileDialog.getSaveFileName(self, "Save configuration to", "", filter="yaml files (*.yaml);;all files (*)")
517
+ savefilename, _ = QtWidgets.QFileDialog.getSaveFileName(
518
+ self, "Save configuration to", "", filter="yaml files (*.yaml);;all files (*)"
519
+ )
507
520
  if len(savefilename):
508
521
  data = dialog.form.get_form_data()
509
522
  logger.info(f" Saving form fields to {savefilename}.")
@@ -514,7 +527,9 @@ class MainWindow(QtWidgets.QMainWindow):
514
527
  def make_cli(arg):
515
528
  script_ext = "cmd" if os.name == "nt" else "sh"
516
529
  savefilename = "train." + script_ext
517
- savefilename, _ = QtWidgets.QFileDialog.getSaveFileName(self, "Script name", savefilename, filter=f"script (*.{script_ext};;all files (*)")
530
+ savefilename, _ = QtWidgets.QFileDialog.getSaveFileName(
531
+ self, "Script name", savefilename, filter=f"script (*.{script_ext};;all files (*)"
532
+ )
518
533
  if len(savefilename):
519
534
  form_data = dialog.form.get_form_data()
520
535
  form_data = _filter_form_data(form_data, is_cli=True)
@@ -696,13 +711,19 @@ class MainWindow(QtWidgets.QMainWindow):
696
711
 
697
712
  params = das.utils.load_params(model_path)
698
713
  if audio.shape[0] < params["nb_hist"]:
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.")
714
+ logger.warning(
715
+ f" Aborting. Audio has fewer samples ({audio.shape[0]}) shorter"
716
+ f" than network chunk size ({params['nb_hist']})."
717
+ " Fix by select longer audio."
718
+ )
700
719
  return
701
720
 
702
721
  # select batch size so that at least 10 batches are run
703
722
  # minimizes loss of annotations from batch size "quantization" errors
704
723
  batch_size = 32
705
- nb_batches = lambda batch_size: int(np.floor((audio.shape[0] - ((batch_size - 1) + params["nb_hist"])) / (params["stride"] * (batch_size))))
724
+ nb_batches = lambda batch_size: int(
725
+ np.floor((audio.shape[0] - ((batch_size - 1) + params["nb_hist"])) / (params["stride"] * (batch_size)))
726
+ )
706
727
  while nb_batches(batch_size) < 10 and batch_size > 1:
707
728
  batch_size -= 1
708
729
 
@@ -765,7 +786,11 @@ class MainWindow(QtWidgets.QMainWindow):
765
786
  # segments['sequence'] = [s for s in segments['sequence'] if s is not None]
766
787
  detected_segment_names = np.unique(segments["sequence"])
767
788
  # if these are indices, get corresponding names
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_:
789
+ if (
790
+ len(detected_segment_names)
791
+ and type(detected_segment_names[0]) is not str
792
+ and type(detected_segment_names[0]) is not np.str_
793
+ ):
769
794
  detected_segment_names = [segments["names"][ii] for ii in detected_segment_names]
770
795
 
771
796
  if len(detected_segment_names) > 0: # and detected_segment_names[0] is not None:
@@ -781,7 +806,9 @@ class MainWindow(QtWidgets.QMainWindow):
781
806
 
782
807
  onsets_seconds = self.ds.sampletime[onsets_samples]
783
808
  offsets_seconds = self.ds.sampletime[offsets_samples]
784
- for name_or_index, onset_seconds, offset_seconds in zip(segments["sequence"], onsets_seconds, offsets_seconds):
809
+ for name_or_index, onset_seconds, offset_seconds in zip(
810
+ segments["sequence"], onsets_seconds, offsets_seconds
811
+ ):
785
812
  if type(name_or_index) is not str and type(detected_segment_names[0]) is not np.str_:
786
813
  segment_name = segments["names"][name_or_index]
787
814
  else:
@@ -830,8 +857,16 @@ class MainWindow(QtWidgets.QMainWindow):
830
857
  except KeyError:
831
858
  pass
832
859
  except KeyError:
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"):
860
+ logger.info(
861
+ f"{filename} no sample rate info in NPZ file."
862
+ f"Need to save 'samplerate' variable with the audio data. Defaulting to {samplerate}"
863
+ )
864
+ elif (
865
+ filename.endswith(".h5")
866
+ or filename.endswith(".hdfs")
867
+ or filename.endswith(".hdf5")
868
+ or filename.endswith(".mat")
869
+ ):
835
870
  # infer data set (for hdf5) and populate form
836
871
  try:
837
872
  # list all data sets in file and add to list
@@ -935,7 +970,9 @@ class MainWindow(QtWidgets.QMainWindow):
935
970
  if not dirname:
936
971
  dirname = QtWidgets.QFileDialog.getExistingDirectory(parent=None, caption="Select data directory")
937
972
  if dirname:
938
- dialog = YamlDialog(yaml_file=package_dir + "/gui/forms/from_dir.yaml", title=f"Dataset from data directory {dirname}")
973
+ dialog = YamlDialog(
974
+ yaml_file=package_dir + "/gui/forms/from_dir.yaml", title=f"Dataset from data directory {dirname}"
975
+ )
939
976
 
940
977
  # initialize form data with cli args
941
978
  dialog.form["pixel_size_mm"] = pixel_size_mm # and un-disable
@@ -1010,7 +1047,9 @@ class MainWindow(QtWidgets.QMainWindow):
1010
1047
 
1011
1048
  # add event categories if they are missing in the dataset
1012
1049
  if "song_events" in ds and "event_categories" not in ds:
1013
- event_categories = ["segment" if "sine" in evt or "syllable" in evt else "event" for evt in ds.event_types.values]
1050
+ event_categories = [
1051
+ "segment" if "sine" in evt or "syllable" in evt else "event" for evt in ds.event_types.values
1052
+ ]
1014
1053
  ds = ds.assign_coords({"event_categories": (("event_types"), event_categories)})
1015
1054
 
1016
1055
  # add missing song types
@@ -1068,7 +1107,9 @@ class MainWindow(QtWidgets.QMainWindow):
1068
1107
  if not filename:
1069
1108
  filename, _ = QtWidgets.QFileDialog.getOpenFileName(parent=None, caption="Select dataset")
1070
1109
  if filename:
1071
- dialog = YamlDialog(yaml_file=package_dir + "/gui/forms/from_zarr.yaml", title=f"Load dataset from zarr file {filename}")
1110
+ dialog = YamlDialog(
1111
+ yaml_file=package_dir + "/gui/forms/from_zarr.yaml", title=f"Load dataset from zarr file {filename}"
1112
+ )
1072
1113
 
1073
1114
  # initialize form data with cli args
1074
1115
  if spec_freq_min is not None:
@@ -1106,7 +1147,9 @@ class MainWindow(QtWidgets.QMainWindow):
1106
1147
 
1107
1148
  # add event categories if they are missing in the dataset
1108
1149
  if "song_events" in ds and "event_categories" not in ds:
1109
- event_categories = ["segment" if "sine" in evt or "syllable" in evt else "event" for evt in ds.event_types.values]
1150
+ event_categories = [
1151
+ "segment" if "sine" in evt or "syllable" in evt else "event" for evt in ds.event_types.values
1152
+ ]
1110
1153
  ds = ds.assign_coords({"event_categories": (("event_types"), event_categories)})
1111
1154
  logger.info(ds)
1112
1155
  vr = None
@@ -1151,11 +1194,15 @@ class MainWindow(QtWidgets.QMainWindow):
1151
1194
 
1152
1195
  def save_dataset(self, qt_keycode=None):
1153
1196
  try:
1154
- savefilename = Path(self.ds.attrs["root"], self.ds.attrs["dat_path"], self.ds.attrs["datename"], f"{self.ds.attrs['datename']}.zarr")
1197
+ savefilename = Path(
1198
+ self.ds.attrs["root"], self.ds.attrs["dat_path"], self.ds.attrs["datename"], f"{self.ds.attrs['datename']}.zarr"
1199
+ )
1155
1200
  except KeyError:
1156
1201
  savefilename = ""
1157
1202
 
1158
- savefilename, _ = QtWidgets.QFileDialog.getSaveFileName(self, "Save dataset to", str(savefilename), filter="zarr files (*.zarr);;all files (*)")
1203
+ savefilename, _ = QtWidgets.QFileDialog.getSaveFileName(
1204
+ self, "Save dataset to", str(savefilename), filter="zarr files (*.zarr);;all files (*)"
1205
+ )
1159
1206
 
1160
1207
  if len(savefilename):
1161
1208
  file_exists = os.path.exists(savefilename)
@@ -1310,7 +1357,7 @@ class PSV(MainWindow):
1310
1357
  self.spec_win = 200
1311
1358
  self.show_songevents = True
1312
1359
  self.movable_events = True
1313
- self.edit_only_current_events = True
1360
+ self.edit_only_current_events = False
1314
1361
  self.show_all_channels = True
1315
1362
  self.select_loudest_channel = False
1316
1363
  self.threshold_mode = False
@@ -1407,10 +1454,16 @@ class PSV(MainWindow):
1407
1454
  self._add_keyed_menuitem(view_video, "Change other fly", self.change_other_fly, "Z")
1408
1455
  self._add_keyed_menuitem(view_video, "Swap flies", self.swap_flies, "X")
1409
1456
  view_video.addSeparator()
1410
- self._add_keyed_menuitem(view_video, "Move poses", partial(self.toggle, "move_poses"), "B", checkable=True, checked=self.move_poses)
1457
+ self._add_keyed_menuitem(
1458
+ view_video, "Move poses", partial(self.toggle, "move_poses"), "B", checkable=True, checked=self.move_poses
1459
+ )
1411
1460
  view_video.addSeparator()
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)
1461
+ self._add_keyed_menuitem(
1462
+ view_video, "Show fly position", partial(self.toggle, "show_dot"), "O", checkable=True, checked=self.show_dot
1463
+ )
1464
+ self._add_keyed_menuitem(
1465
+ view_video, "Show poses", partial(self.toggle, "show_poses"), "P", checkable=True, checked=self.show_poses
1466
+ )
1414
1467
 
1415
1468
  view_audio = self.bar.addMenu("Audio")
1416
1469
  self._add_keyed_menuitem(view_audio, "Play waveform through speakers", self.play_audio, "E")
@@ -1434,7 +1487,9 @@ class PSV(MainWindow):
1434
1487
  self._add_keyed_menuitem(view_audio, "Select previous channel", self.set_next_channel, "Up")
1435
1488
  self._add_keyed_menuitem(view_audio, "Select next channel", self.set_prev_channel, "Down")
1436
1489
  view_audio.addSeparator()
1437
- self._add_keyed_menuitem(view_audio, "Show spectrogram", partial(self.toggle, "show_spec"), None, checkable=True, checked=self.show_spec)
1490
+ self._add_keyed_menuitem(
1491
+ view_audio, "Show spectrogram", partial(self.toggle, "show_spec"), None, checkable=True, checked=self.show_spec
1492
+ )
1438
1493
  self._add_keyed_menuitem(view_audio, "Increase frequency resolution", self.inc_freq_res, "R")
1439
1494
  self._add_keyed_menuitem(view_audio, "Increase temporal resolution", self.dec_freq_res, "T")
1440
1495
  view_audio.addSeparator()
@@ -1488,20 +1543,34 @@ class PSV(MainWindow):
1488
1543
  self._add_keyed_menuitem(view_annotations, "Generate proposal by envelope thresholding", self.threshold, "I")
1489
1544
  self._add_keyed_menuitem(view_annotations, "Adjust thresholding mode", self.set_envelope_computation)
1490
1545
  view_annotations.addSeparator()
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")
1546
+ self._add_keyed_menuitem(
1547
+ view_annotations, "Approve proposals for active song type in view", self.approve_active_proposals, "G"
1548
+ )
1549
+ self._add_keyed_menuitem(
1550
+ view_annotations, "Approve proposals for all song types in view", self.approve_all_proposals, "H"
1551
+ )
1493
1552
 
1494
1553
  view_view = self.bar.addMenu("View")
1495
1554
  self._add_keyed_menuitem(view_view, "Video, waveform, and spectrogram display parameters", self.set_spec_freq)
1496
1555
  view_view.addSeparator()
1497
1556
  # TODO? only show these if tracks and/or video
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)
1557
+ self._add_keyed_menuitem(
1558
+ view_view, "Show spectrogram", partial(self.toggle, "show_spec"), None, checkable=True, checked=self.show_spec
1559
+ )
1560
+ self._add_keyed_menuitem(
1561
+ view_view, "Show waveform", partial(self.toggle, "show_trace"), None, checkable=True, checked=self.show_trace
1562
+ )
1563
+ self._add_keyed_menuitem(
1564
+ view_view, "Show ethogram", partial(self.toggle, "show_annot"), None, checkable=True, checked=self.show_annot
1565
+ )
1501
1566
  if "pose_positions_allo" in self.ds:
1502
- self._add_keyed_menuitem(view_view, "Show tracks", partial(self.toggle, "show_tracks"), None, checkable=True, checked=self.show_tracks)
1567
+ self._add_keyed_menuitem(
1568
+ view_view, "Show tracks", partial(self.toggle, "show_tracks"), None, checkable=True, checked=self.show_tracks
1569
+ )
1503
1570
  if self.vr is not None:
1504
- self._add_keyed_menuitem(view_view, "Show movie", partial(self.toggle, "show_movie"), None, checkable=True, checked=self.show_movie)
1571
+ self._add_keyed_menuitem(
1572
+ view_view, "Show movie", partial(self.toggle, "show_movie"), None, checkable=True, checked=self.show_movie
1573
+ )
1505
1574
 
1506
1575
  self.hl = QtWidgets.QHBoxLayout()
1507
1576
 
@@ -1872,7 +1941,9 @@ class PSV(MainWindow):
1872
1941
 
1873
1942
  def delete_current_events(self, qt_keycode):
1874
1943
  if self.current_event_index is not None:
1875
- deleted_events = self.event_times.delete_range(self.current_event_name, self.time0 / self.fs_song, self.time1 / self.fs_song)
1944
+ deleted_events = self.event_times.delete_range(
1945
+ self.current_event_name, self.time0 / self.fs_song, self.time1 / self.fs_song
1946
+ )
1876
1947
  nb_deleted_events = len(deleted_events)
1877
1948
  if nb_deleted_events:
1878
1949
  logger.info(f" Deleted {nb_deleted_events} annotation(s) of type {self.current_event_name}.")
@@ -1894,7 +1965,9 @@ class PSV(MainWindow):
1894
1965
  def threshold(self, qt_keycode):
1895
1966
  if self.STOP and self.current_event_name is not None:
1896
1967
  if self.event_times.categories[self.current_event_name] == "event":
1897
- indexes = peakutils.indexes(self.envelope, thres=self.slice_view.threshold, min_dist=self.thres_min_dist * self.fs_song, thres_abs=True)
1968
+ indexes = peakutils.indexes(
1969
+ self.envelope, thres=self.slice_view.threshold, min_dist=self.thres_min_dist * self.fs_song, thres_abs=True
1970
+ )
1898
1971
  # add events to current song type
1899
1972
  for t in self.x[indexes]:
1900
1973
  self.event_times.add_time(self.current_event_name, t)
@@ -2044,7 +2117,9 @@ class PSV(MainWindow):
2044
2117
  dialog.exec_()
2045
2118
 
2046
2119
  def set_envelope_computation(self, qt_keycode):
2047
- dialog = YamlDialog(yaml_file=package_dir + "/gui/forms/envelope_computation.yaml", title="Set options for envelope computation")
2120
+ dialog = YamlDialog(
2121
+ yaml_file=package_dir + "/gui/forms/envelope_computation.yaml", title="Set options for envelope computation"
2122
+ )
2048
2123
 
2049
2124
  dialog.form["thres_min_dist"] = self.thres_min_dist
2050
2125
  dialog.form["thres_env_std"] = self.thres_env_std
@@ -2122,7 +2197,9 @@ class PSV(MainWindow):
2122
2197
  i1 = int(self.time1 / self.fs_ratio)
2123
2198
 
2124
2199
  self.x_tracks = self.ds.time.data[i0:i1]
2125
- self.y_tracks = self.ds.pose_positions_allo.data[i0:i1, self.focal_fly, self.track_sel_names, self.track_sel_coords]
2200
+ self.y_tracks = self.ds.pose_positions_allo.data[
2201
+ i0:i1, self.focal_fly, self.track_sel_names, self.track_sel_coords
2202
+ ]
2126
2203
  self.tracks_view.update_trace()
2127
2204
  self.tracks_view.show()
2128
2205
  else:
@@ -2169,9 +2246,13 @@ class PSV(MainWindow):
2169
2246
  if self.event_times.categories[event_name] == "segment":
2170
2247
  for onset, offset in zip(events_in_view[:, 0], events_in_view[:, 1]):
2171
2248
  if self.show_trace:
2172
- self.slice_view.add_segment(onset, offset, event_index, brush=event_brush, pen=event_pen, movable=movable, text=segment_text)
2249
+ self.slice_view.add_segment(
2250
+ onset, offset, event_index, brush=event_brush, pen=event_pen, movable=movable, text=segment_text
2251
+ )
2173
2252
  if self.show_tracks:
2174
- self.tracks_view.add_segment(onset, offset, event_index, brush=event_brush, pen=event_pen, movable=movable, text=segment_text)
2253
+ self.tracks_view.add_segment(
2254
+ onset, offset, event_index, brush=event_brush, pen=event_pen, movable=movable, text=segment_text
2255
+ )
2175
2256
  if self.show_annot:
2176
2257
  self.annot_view.add_segment(
2177
2258
  onset,
@@ -2183,7 +2264,9 @@ class PSV(MainWindow):
2183
2264
  text=segment_text,
2184
2265
  )
2185
2266
  if self.show_spec:
2186
- self.spec_view.add_segment(onset, offset, event_index, brush=event_brush, pen=event_pen, movable=movable, text=segment_text)
2267
+ self.spec_view.add_segment(
2268
+ onset, offset, event_index, brush=event_brush, pen=event_pen, movable=movable, text=segment_text
2269
+ )
2187
2270
  elif self.event_times.categories[event_name] == "event":
2188
2271
  if self.show_trace:
2189
2272
  self.slice_view.add_event(events_in_view[:, 0], event_index, event_pen, movable=movable, text=segment_text)
@@ -2216,15 +2299,19 @@ class PSV(MainWindow):
2216
2299
 
2217
2300
  new_region = region.getRegion()
2218
2301
  self.event_times.move_time(event_name_to_move, region.bounds, new_region)
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.")
2302
+ logger.info(
2303
+ 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()
2222
- if mp > 0 and mp < 1:
2223
- new_event_idx = int(mp * self.nb_eventtypes)
2224
- new_event_name = self.event_times.names[new_event_idx]
2225
- _, old_name, new_name = self.event_times.change_name(new_region[0], new_event_name)
2226
- if old_name is not None:
2227
- logger.info(f" Changed from {old_name} to {new_name}.")
2306
+ # FIXME for moving annotations in ethogram - fails in pyside6
2307
+ if self.annot_view.mousePoint is not None:
2308
+ mp = self.annot_view.mousePoint.y()
2309
+ if mp > 0 and mp < 1:
2310
+ new_event_idx = int(mp * self.nb_eventtypes)
2311
+ new_event_name = self.event_times.names[new_event_idx]
2312
+ _, old_name, new_name = self.event_times.change_name(new_region[0], new_event_name)
2313
+ if old_name is not None:
2314
+ logger.info(f" Changed from {old_name} to {new_name}.")
2228
2315
 
2229
2316
  self.update_xy()
2230
2317
 
@@ -2235,19 +2322,20 @@ class PSV(MainWindow):
2235
2322
  event_name_to_move = self.current_event_name
2236
2323
  if self.current_event_index != position.event_index:
2237
2324
  event_name_to_move = self.event_times.names[position.event_index]
2238
-
2325
+ print(position.pos(), position.position)
2239
2326
  new_position = position.pos()[0]
2240
2327
  self.event_times.move_time(event_name_to_move, position.position, new_position)
2241
2328
  logger.info(f" Moved {event_name_to_move} from t={position.position:1.4f} to {new_position:1.4f} seconds.")
2242
2329
 
2243
- mp = self.annot_view.mousePoint.y()
2244
- if mp > 0 and mp < 1:
2245
- new_event_idx = int(mp * self.nb_eventtypes)
2246
- new_event_name = self.event_times.names[new_event_idx]
2247
- print(new_event_name)
2248
- _, old_name, new_name = self.event_times.change_name(new_position, new_event_name, tol=1)
2249
- if old_name is not None:
2250
- logger.info(f" Changed from {old_name} to {new_name}.")
2330
+ # FIXME for moving annotations in ethogram - fails in pyside6
2331
+ if self.annot_view.mousePoint is not None:
2332
+ mp = self.annot_view.mousePoint.y()
2333
+ if mp > 0 and mp < 1:
2334
+ new_event_idx = int(mp * self.nb_eventtypes)
2335
+ new_event_name = self.event_times.names[new_event_idx]
2336
+ _, old_name, new_name = self.event_times.change_name(new_position, new_event_name, old_name=event_name_to_move)
2337
+ if old_name is not None:
2338
+ logger.info(f" Changed from {old_name} to {new_name}.")
2251
2339
 
2252
2340
  self.update_xy()
2253
2341
 
@@ -2285,7 +2373,10 @@ class PSV(MainWindow):
2285
2373
  fly_pos = self.ds.pose_positions_allo.data[self.index_other, :, self.pose_center_index, :]
2286
2374
  fly_pos = np.array(fly_pos) # in case this is a dask.array
2287
2375
  if self.crop: # transform fly pos to coordinates of the cropped box
2288
- box_center = self.ds.pose_positions_allo.data[self.index_other, self.focal_fly, self.pose_center_index] + self.box_size / 2
2376
+ box_center = (
2377
+ self.ds.pose_positions_allo.data[self.index_other, self.focal_fly, self.pose_center_index]
2378
+ + self.box_size / 2
2379
+ )
2289
2380
  box_center = np.array(box_center) # in case this is a dask.array
2290
2381
  fly_pos = fly_pos - box_center
2291
2382
  fly_dist = np.sum((fly_pos - np.array([mouseY, mouseX])) ** 2, axis=-1)
@@ -2325,7 +2416,9 @@ class PSV(MainWindow):
2325
2416
  if self.event_times.categories[self.current_event_name] == "event":
2326
2417
  logger.info(f" Changed event at {changed_time[0]:1.4f} from {old_name} to {new_name}.")
2327
2418
  else:
2328
- logger.info(f" Changed segment at {changed_time[0]:1.4f}:{changed_time[1]:1.4f} from {old_name} to {new_name}.")
2419
+ logger.info(
2420
+ f" Changed segment at {changed_time[0]:1.4f}:{changed_time[1]:1.4f} from {old_name} to {new_name}."
2421
+ )
2329
2422
  self.update_xy()
2330
2423
  elif mouseButton == 1: # add event
2331
2424
  if self.current_event_index is not None:
@@ -2345,12 +2438,16 @@ class PSV(MainWindow):
2345
2438
  stop_seconds=mouseT,
2346
2439
  channel=self.current_channel_index,
2347
2440
  )
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.")
2441
+ logger.info(
2442
+ f" Added {self.current_event_name} on channel {self.current_channel_index} at t=[{self.sinet0:1.4f}:{mouseT:1.4f}] seconds."
2443
+ )
2349
2444
  self.sinet0 = None
2350
2445
  if self.event_times.categories[self.current_event_name] == "event":
2351
2446
  self.sinet0 = None
2352
2447
  self.event_times.add_time(self.current_event_name, start_seconds=mouseT, channel=self.current_channel_index)
2353
- logger.info(f" Added {self.current_event_name} on channel {self.current_channel_index} at t={mouseT:1.4f} seconds.")
2448
+ logger.info(
2449
+ f" Added {self.current_event_name} on channel {self.current_channel_index} at t={mouseT:1.4f} seconds."
2450
+ )
2354
2451
  self.update_xy()
2355
2452
  else:
2356
2453
  self.sinet0 = None
@@ -394,7 +394,9 @@ def make(
394
394
  store.attrs["data_splits"] = data_split_dict
395
395
  logger.info("Done.")
396
396
  # report
397
- logger.info(f" Got {store['train']['x'].shape}, {store['val']['x'].shape}, {store['test']['x'].shape} train/val/test samples.")
397
+ logger.info(
398
+ f" Got {store['train']['x'].shape}, {store['val']['x'].shape}, {store['test']['x'].shape} train/val/test samples."
399
+ )
398
400
  if to_npy_dir: # save as npy_dir
399
401
  logger.info(f" Saving to {store_folder}.")
400
402
  das.npy_dir.save(store_folder, store)
@@ -409,7 +411,9 @@ def make(
409
411
  for key in store.keys():
410
412
  ann = event_utils.traces_to_eventtimes(store[key]["y"][:, 1:].T, names, types, events_are_binary=False)
411
413
  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
414
+ if (
415
+ v.ndim == 1
416
+ ): # add start/stop seconds for events - traces_to_eventtimes only returns event start times but annot expects both start and end
413
417
  v = np.stack((v, v)).T
414
418
  v = v / fs # convert indices to seconds
415
419
  ann[k] = v
@@ -26,6 +26,7 @@ is :py:method:`FormBuilderLayout.add_item()`. Look there if you want to know
26
26
  what's supported in the YAML file, what exactly each field type does, or if you
27
27
  want to add a new type of supported form field.
28
28
  """
29
+
29
30
  # modified from https://sleap.ai/_modules/sleap/gui/formbuilder.html
30
31
  import yaml
31
32
  from typing import Any, Dict, List, Optional, Union, Text
@@ -572,8 +572,6 @@ class AnnotView(pg.PlotWidget):
572
572
  # additionally make names of trace and event arrays in ds args?
573
573
  super().__init__()
574
574
  self.setMouseEnabled(x=False, y=False)
575
- # this should be just a link/ref so changes in ds made by the controller will propagate
576
- # mabe make Model as thin wrapper around ds that also handles ion and use ref to Modle instance
577
575
  self.disableAutoRange()
578
576
  self.enableAutoRange(False, False)
579
577
  self.setDefaultPadding(0.0)
@@ -655,12 +653,20 @@ class AnnotView(pg.PlotWidget):
655
653
  return np.interp(pos, self.xrange, self.m.trange)
656
654
 
657
655
  def _click(self, event):
658
- event.accept()
656
+ # event.accept()
659
657
  pos = event.pos()
660
658
  mouseT = self.getPlotItem().getViewBox().mapSceneToView(pos).x()
661
659
  self.callback(mouseT, event.button())
662
660
 
661
+ # def mouseMoveEvent(self, ev):
662
+ # pos = ev.pos()
663
+ # # print(pos)
664
+ # if self.sceneBoundingRect().contains(pos):
665
+ # self.mousePoint = self.getPlotItem().vb.mapSceneToView(pos)
666
+ # print(self.mousePoint)
667
+
663
668
  def mouseMoved(self, pos):
669
+ # pass
664
670
  if self.sceneBoundingRect().contains(pos):
665
671
  self.mousePoint = self.getPlotItem().vb.mapSceneToView(pos)
666
672
 
@@ -5,6 +5,7 @@ should return:
5
5
  non_audio_data: np.array[time, samples]
6
6
  samplerate: Optional[float]
7
7
  """
8
+
8
9
  # [x] daq.h5
9
10
  # [x] wav, ....
10
11
  # [x] npz, npy
@@ -3,6 +3,7 @@
3
3
  should return:
4
4
  x: pd.DataFrame[frames, (variables)]
5
5
  """
6
+
6
7
  import numpy as np
7
8
  import pandas as pd
8
9
  from .. import io
@@ -1,6 +1,7 @@
1
1
  """Load parameters for DLP movies
2
2
 
3
3
  """
4
+
4
5
  import numpy as np
5
6
  import pandas as pd
6
7
  import h5py
@@ -7,6 +7,7 @@ should return:
7
7
  first_pose_frame: int
8
8
  last_pose_frame: int
9
9
  """
10
+
10
11
  import h5py
11
12
  import numpy as np
12
13
  import xarray as xr
@@ -292,8 +293,12 @@ class Sleap(Poses, io.BaseProvider):
292
293
 
293
294
  # indices and coordinate logic for sleap tracks. This could be cleaned/simplified later
294
295
  nb_flies = tracks.shape[0]
295
- thorax_idx = np.argwhere(pose_parts == b"thorax")[0][0]
296
- head_idx = np.argwhere(pose_parts == b"head")[0][0]
296
+ try:
297
+ thorax_idx = np.argwhere(pose_parts == b"thorax")[0][0]
298
+ head_idx = np.argwhere(pose_parts == b"head")[0][0]
299
+ except IndexError:
300
+ thorax_idx = 0
301
+ head_idx = 1
297
302
  x_idx = 0
298
303
  y_idx = 1
299
304
 
@@ -1,4 +1,5 @@
1
1
  """Load files created by various analysis programs created."""
2
+
2
3
  import numpy as np
3
4
  import pandas as pd
4
5
  import scipy.interpolate
@@ -1,4 +1,5 @@
1
1
  """Calculate metrics from behavioral data."""
2
+
2
3
  import numpy as np
3
4
  import scipy.signal.windows
4
5
 
@@ -1,4 +1,5 @@
1
1
  """Create self-documenting xarray dataset from behavioral recordings and annotations."""
2
+
2
3
  import numpy as np
3
4
  from samplestamps.samplestamps import SampStamp, SimpleStamp
4
5
  import scipy.interpolate
File without changes
File without changes
File without changes