prpy 0.3.2__tar.gz → 0.3.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 (60) hide show
  1. {prpy-0.3.2/prpy.egg-info → prpy-0.3.4}/PKG-INFO +1 -1
  2. {prpy-0.3.2 → prpy-0.3.4}/prpy/numpy/image.py +2 -1
  3. {prpy-0.3.2 → prpy-0.3.4}/prpy/numpy/metric.py +71 -34
  4. {prpy-0.3.2 → prpy-0.3.4}/prpy/numpy/physio.py +84 -13
  5. {prpy-0.3.2 → prpy-0.3.4/prpy.egg-info}/PKG-INFO +1 -1
  6. {prpy-0.3.2 → prpy-0.3.4}/tests/test_numpy_metric.py +44 -11
  7. {prpy-0.3.2 → prpy-0.3.4}/tests/test_numpy_physio.py +11 -8
  8. {prpy-0.3.2 → prpy-0.3.4}/.gitignore +0 -0
  9. {prpy-0.3.2 → prpy-0.3.4}/LICENSE +0 -0
  10. {prpy-0.3.2 → prpy-0.3.4}/MANIFEST.in +0 -0
  11. {prpy-0.3.2 → prpy-0.3.4}/README.md +0 -0
  12. {prpy-0.3.2 → prpy-0.3.4}/prpy/__init__.py +0 -0
  13. {prpy-0.3.2 → prpy-0.3.4}/prpy/constants.py +0 -0
  14. {prpy-0.3.2 → prpy-0.3.4}/prpy/ffmpeg/__init__.py +0 -0
  15. {prpy-0.3.2 → prpy-0.3.4}/prpy/ffmpeg/probe.py +0 -0
  16. {prpy-0.3.2 → prpy-0.3.4}/prpy/ffmpeg/readwrite.py +0 -0
  17. {prpy-0.3.2 → prpy-0.3.4}/prpy/ffmpeg/utils.py +0 -0
  18. {prpy-0.3.2 → prpy-0.3.4}/prpy/helpers.py +0 -0
  19. {prpy-0.3.2 → prpy-0.3.4}/prpy/numpy/__init__.py +0 -0
  20. {prpy-0.3.2 → prpy-0.3.4}/prpy/numpy/core.py +0 -0
  21. {prpy-0.3.2 → prpy-0.3.4}/prpy/numpy/detect.py +0 -0
  22. {prpy-0.3.2 → prpy-0.3.4}/prpy/numpy/face.py +0 -0
  23. {prpy-0.3.2 → prpy-0.3.4}/prpy/numpy/filters.py +0 -0
  24. {prpy-0.3.2 → prpy-0.3.4}/prpy/numpy/freq.py +0 -0
  25. {prpy-0.3.2 → prpy-0.3.4}/prpy/numpy/image_ops.c +0 -0
  26. {prpy-0.3.2 → prpy-0.3.4}/prpy/numpy/interp.py +0 -0
  27. {prpy-0.3.2 → prpy-0.3.4}/prpy/numpy/rolling.py +0 -0
  28. {prpy-0.3.2 → prpy-0.3.4}/prpy/numpy/stride_tricks.py +0 -0
  29. {prpy-0.3.2 → prpy-0.3.4}/prpy/numpy/utils.py +0 -0
  30. {prpy-0.3.2 → prpy-0.3.4}/prpy/tensorflow/__init__.py +0 -0
  31. {prpy-0.3.2 → prpy-0.3.4}/prpy/tensorflow/image.py +0 -0
  32. {prpy-0.3.2 → prpy-0.3.4}/prpy/tensorflow/loss.py +0 -0
  33. {prpy-0.3.2 → prpy-0.3.4}/prpy/tensorflow/lr_schedule.py +0 -0
  34. {prpy-0.3.2 → prpy-0.3.4}/prpy/tensorflow/model_saver.py +0 -0
  35. {prpy-0.3.2 → prpy-0.3.4}/prpy/tensorflow/nan.py +0 -0
  36. {prpy-0.3.2 → prpy-0.3.4}/prpy/tensorflow/optimizer.py +0 -0
  37. {prpy-0.3.2 → prpy-0.3.4}/prpy/tensorflow/signal.py +0 -0
  38. {prpy-0.3.2 → prpy-0.3.4}/prpy/torch/__init__.py +0 -0
  39. {prpy-0.3.2 → prpy-0.3.4}/prpy/torch/model_saver.py +0 -0
  40. {prpy-0.3.2 → prpy-0.3.4}/prpy.egg-info/SOURCES.txt +0 -0
  41. {prpy-0.3.2 → prpy-0.3.4}/prpy.egg-info/dependency_links.txt +0 -0
  42. {prpy-0.3.2 → prpy-0.3.4}/prpy.egg-info/requires.txt +0 -0
  43. {prpy-0.3.2 → prpy-0.3.4}/prpy.egg-info/top_level.txt +0 -0
  44. {prpy-0.3.2 → prpy-0.3.4}/pyproject.toml +0 -0
  45. {prpy-0.3.2 → prpy-0.3.4}/setup.cfg +0 -0
  46. {prpy-0.3.2 → prpy-0.3.4}/setup.py +0 -0
  47. {prpy-0.3.2 → prpy-0.3.4}/tests/conftest.py +0 -0
  48. {prpy-0.3.2 → prpy-0.3.4}/tests/test_ffmpeg.py +0 -0
  49. {prpy-0.3.2 → prpy-0.3.4}/tests/test_numpy_core.py +0 -0
  50. {prpy-0.3.2 → prpy-0.3.4}/tests/test_numpy_detect.py +0 -0
  51. {prpy-0.3.2 → prpy-0.3.4}/tests/test_numpy_face.py +0 -0
  52. {prpy-0.3.2 → prpy-0.3.4}/tests/test_numpy_filters.py +0 -0
  53. {prpy-0.3.2 → prpy-0.3.4}/tests/test_numpy_freq.py +0 -0
  54. {prpy-0.3.2 → prpy-0.3.4}/tests/test_numpy_image.py +0 -0
  55. {prpy-0.3.2 → prpy-0.3.4}/tests/test_numpy_interp.py +0 -0
  56. {prpy-0.3.2 → prpy-0.3.4}/tests/test_numpy_rolling.py +0 -0
  57. {prpy-0.3.2 → prpy-0.3.4}/tests/test_numpy_stride_tricks.py +0 -0
  58. {prpy-0.3.2 → prpy-0.3.4}/tests/test_numpy_utils.py +0 -0
  59. {prpy-0.3.2 → prpy-0.3.4}/tests/test_tensorflow.py +0 -0
  60. {prpy-0.3.2 → prpy-0.3.4}/tests/test_torch.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: prpy
3
- Version: 0.3.2
3
+ Version: 0.3.4
4
4
  Summary: Collection of Python utils for signal, image, and video processing
5
5
  Author-email: Philipp Rouast <philipp@rouast.com>
6
6
  License: MIT License
@@ -22,7 +22,6 @@ import imghdr
22
22
  import math
23
23
  import numpy as np
24
24
  import os
25
- from PIL import Image
26
25
  import logging
27
26
  from typing import Union, Tuple
28
27
 
@@ -244,6 +243,7 @@ def probe_image_inputs(
244
243
  # Image
245
244
  if not allow_image:
246
245
  raise ValueError(f"allow_image={allow_image}, but received a path to an image file.")
246
+ from PIL import Image
247
247
  with Image.open(inputs) as img:
248
248
  width, height = img.size
249
249
  channels = len(img.getbands())
@@ -363,6 +363,7 @@ def parse_image_inputs(
363
363
  if not allow_image:
364
364
  raise ValueError(f"allow_image={allow_image}, but received a path to an image file.")
365
365
  try:
366
+ from PIL import Image
366
367
  inputs = np.asarray(Image.open(open(inputs, 'rb')))
367
368
  except Exception as e:
368
369
  raise ValueError(f"Problem reading image from {inputs}: {e}")
@@ -1,4 +1,4 @@
1
- # Copyright (c) 2024 Philipp Rouast
1
+ # Copyright (c) 2025 Philipp Rouast
2
2
  #
3
3
  # Permission is hereby granted, free of charge, to any person obtaining a copy
4
4
  # of this software and associated documentation files (the "Software"), to deal
@@ -22,7 +22,15 @@ import numpy as np
22
22
  from scipy import signal
23
23
  from typing import Union
24
24
 
25
- def mag2db(mag: Union[np.ndarray, np.float64]) -> np.ndarray:
25
+ def _masked_sum(x, mask, axis, *, keepdims=False):
26
+ return np.where(mask, x, 0.0).sum(axis=axis, keepdims=keepdims)
27
+
28
+ def _safe_mean(x, mask, axis):
29
+ n_valid = mask.sum(axis=axis, keepdims=True)
30
+ n_safe = np.where(n_valid == 0, 1, n_valid)
31
+ return _masked_sum(x, mask, axis, keepdims=True) / n_safe
32
+
33
+ def _mag2db(mag: Union[np.ndarray, np.float64]) -> np.ndarray:
26
34
  """
27
35
  Magnitude to decibels element-wise.
28
36
 
@@ -37,7 +45,9 @@ def mag2db(mag: Union[np.ndarray, np.float64]) -> np.ndarray:
37
45
  def mae(
38
46
  y_true: np.ndarray,
39
47
  y_pred: np.ndarray,
40
- axis: Union[int, None] = -1
48
+ axis: Union[int, None] = -1,
49
+ *,
50
+ ignore_nan: bool = False
41
51
  ) -> np.ndarray:
42
52
  """
43
53
  Mean absolute error
@@ -46,6 +56,7 @@ def mae(
46
56
  y_true: True values. Shape (..., dim_n, dim_axis)
47
57
  y_pred: Predicted values. Shape (..., dim_n, dim_axis)
48
58
  axis: Axis along which the means are computed
59
+ ignore_nan: If true, metric is computed on non-NaN elements only.
49
60
  Returns:
50
61
  mae: The mean absolute error. Shape (..., dim_n)
51
62
  """
@@ -53,15 +64,23 @@ def mae(
53
64
  y_true = np.asarray(y_true)
54
65
  y_pred = np.asarray(y_pred)
55
66
  assert y_true.shape == y_pred.shape
56
- if np.all(np.isfinite(y_true)) and np.all(np.isfinite(y_pred)):
57
- return np.mean(np.abs(y_true - y_pred), axis=axis)
58
- else:
59
- return np.full(y_true.shape[:axis], fill_value=np.nan, dtype=y_true.dtype)
67
+ if not ignore_nan:
68
+ if np.all(np.isfinite(y_true)) and np.all(np.isfinite(y_pred)):
69
+ return np.mean(np.abs(y_true - y_pred), axis=axis)
70
+ else:
71
+ return np.full(y_true.shape[:axis], fill_value=np.nan, dtype=y_true.dtype)
72
+ mask = np.isfinite(y_true) & np.isfinite(y_pred)
73
+ err = np.abs(y_true - y_pred)
74
+ n_valid = mask.sum(axis=axis)
75
+ err_sum = _masked_sum(err, mask, axis)
76
+ return np.where(n_valid > 0, err_sum / n_valid, np.nan)
60
77
 
61
78
  def mse(
62
79
  y_true: np.ndarray,
63
80
  y_pred: np.ndarray,
64
- axis: Union[int, None] = -1
81
+ axis: Union[int, None] = -1,
82
+ *,
83
+ ignore_nan: bool = False
65
84
  ) -> np.ndarray:
66
85
  """
67
86
  Mean squared error
@@ -70,6 +89,7 @@ def mse(
70
89
  y_true: True values. Shape (..., dim_n, dim_axis)
71
90
  y_pred: Predicted values. Shape (..., dim_n, dim_axis)
72
91
  axis: Axis along which the means are computed
92
+ ignore_nan: If true, metric is computed on non-NaN elements only.
73
93
  Returns:
74
94
  mse: The mean squared error. Shape (..., dim_n)
75
95
  """
@@ -77,15 +97,23 @@ def mse(
77
97
  y_true = np.asarray(y_true)
78
98
  y_pred = np.asarray(y_pred)
79
99
  assert y_true.shape == y_pred.shape
80
- if np.all(np.isfinite(y_true)) and np.all(np.isfinite(y_pred)):
81
- return np.mean(np.square(y_true - y_pred), axis=axis)
82
- else:
83
- return np.full(y_true.shape[:axis], fill_value=np.nan, dtype=y_true.dtype)
100
+ if not ignore_nan:
101
+ if np.all(np.isfinite(y_true)) and np.all(np.isfinite(y_pred)):
102
+ return np.mean(np.square(y_true - y_pred), axis=axis)
103
+ else:
104
+ return np.full(y_true.shape[:axis], fill_value=np.nan, dtype=y_true.dtype)
105
+ mask = np.isfinite(y_true) & np.isfinite(y_pred)
106
+ sq = (y_true - y_pred) ** 2
107
+ n_valid = mask.sum(axis=axis)
108
+ sq_sum = _masked_sum(sq, mask, axis)
109
+ return np.where(n_valid > 0, sq_sum / n_valid, np.nan)
84
110
 
85
111
  def rmse(
86
112
  y_true: np.ndarray,
87
113
  y_pred: np.ndarray,
88
- axis: Union[int, None] = -1
114
+ axis: Union[int, None] = -1,
115
+ *,
116
+ ignore_nan: bool = False
89
117
  ) -> np.ndarray:
90
118
  """
91
119
  Root mean squared error
@@ -94,22 +122,18 @@ def rmse(
94
122
  y_true: True values. Shape (..., dim_n, dim_axis)
95
123
  y_pred: Predicted values. Shape (..., dim_n, dim_axis)
96
124
  axis: Axis along which the means are computed
125
+ ignore_nan: If true, metric is computed on non-NaN elements only.
97
126
  Returns:
98
127
  rmse: The root mean squared error. Shape (..., dim_n)
99
128
  """
100
- assert axis is None or isinstance(axis, int)
101
- y_true = np.asarray(y_true)
102
- y_pred = np.asarray(y_pred)
103
- assert y_true.shape == y_pred.shape
104
- if np.all(np.isfinite(y_true)) and np.all(np.isfinite(y_pred)):
105
- return np.sqrt(np.mean(np.square(y_true - y_pred), axis=axis))
106
- else:
107
- return np.full(y_true.shape[:axis], fill_value=np.nan, dtype=y_true.dtype)
129
+ return np.sqrt(mse(y_true, y_pred, axis=axis, ignore_nan=ignore_nan))
108
130
 
109
131
  def cor(
110
132
  y_true: np.ndarray,
111
133
  y_pred: np.ndarray,
112
- axis: Union[int, None] = -1
134
+ axis: Union[int, None] = -1,
135
+ *,
136
+ ignore_nan: bool = False,
113
137
  ) -> np.ndarray:
114
138
  """
115
139
  Pearson's correlation coefficient
@@ -118,6 +142,7 @@ def cor(
118
142
  y_true: True values. Shape (..., dim_n, dim_axis)
119
143
  y_pred: Predicted values. Shape (..., dim_n, dim_axis)
120
144
  axis: Axis along which correlations are computed
145
+ ignore_nan: If true, metric is computed on non-NaN elements only.
121
146
  Returns:
122
147
  cor: The correlation coefficients. Shape (..., dim_n)
123
148
  """
@@ -125,17 +150,29 @@ def cor(
125
150
  y_true = np.asarray(y_true)
126
151
  y_pred = np.asarray(y_pred)
127
152
  assert y_true.shape == y_pred.shape
128
- if np.all(np.isfinite(y_true)) and np.all(np.isfinite(y_pred)):
129
- res_true = y_true - np.mean(y_true, axis=axis, keepdims=True)
130
- res_pred = y_pred - np.mean(y_pred, axis=axis, keepdims=True)
131
- cov = np.mean(res_true * res_pred, axis=axis)
132
- var_true = np.mean(res_true**2, axis=axis)
133
- var_pred = np.mean(res_pred**2, axis=axis)
134
- sigma_true = np.sqrt(var_true)
135
- sigma_pred = np.sqrt(var_pred)
136
- return cov / (sigma_true * sigma_pred)
137
- else:
138
- return np.full(y_true.shape[:axis], fill_value=np.nan, dtype=y_true.dtype)
153
+ if not ignore_nan:
154
+ if np.all(np.isfinite(y_true)) and np.all(np.isfinite(y_pred)):
155
+ res_true = y_true - np.mean(y_true, axis=axis, keepdims=True)
156
+ res_pred = y_pred - np.mean(y_pred, axis=axis, keepdims=True)
157
+ cov = np.mean(res_true * res_pred, axis=axis)
158
+ var_true = np.mean(res_true**2, axis=axis)
159
+ var_pred = np.mean(res_pred**2, axis=axis)
160
+ sigma_true = np.sqrt(var_true)
161
+ sigma_pred = np.sqrt(var_pred)
162
+ return cov / (sigma_true * sigma_pred)
163
+ else:
164
+ return np.full(y_true.shape[:axis], fill_value=np.nan, dtype=y_true.dtype)
165
+ mask = np.isfinite(y_true) & np.isfinite(y_pred)
166
+ mean_t = _safe_mean(y_true, mask, axis)
167
+ mean_p = _safe_mean(y_pred, mask, axis)
168
+ r_true = np.where(mask, y_true - mean_t, 0.0)
169
+ r_pred = np.where(mask, y_pred - mean_p, 0.0)
170
+ n_valid = mask.sum(axis=axis)
171
+ cov = (r_true * r_pred).sum(axis=axis) / np.where(n_valid == 0, 1, n_valid)
172
+ var_t = (r_true ** 2).sum(axis=axis) / np.where(n_valid == 0, 1, n_valid)
173
+ var_p = (r_pred ** 2).sum(axis=axis) / np.where(n_valid == 0, 1, n_valid)
174
+ corr = cov / (np.sqrt(var_t) * np.sqrt(var_p))
175
+ return np.where(n_valid > 0, corr, np.nan)
139
176
 
140
177
  def snr(
141
178
  f_true: np.ndarray,
@@ -186,7 +223,7 @@ def snr(
186
223
  s_power = np.sum(pxx * gt_mask, axis=-1)
187
224
  f_mask = (f >= f_min) & (f <= f_max)
188
225
  all_power = np.sum(pxx * f_mask, axis=-1)
189
- snr = mag2db(s_power / (all_power - s_power))
226
+ snr = _mag2db(s_power / (all_power - s_power))
190
227
  return np.squeeze(snr)
191
228
  else:
192
229
  return np.full(y_pred.shape[:-1], fill_value=np.nan, dtype=y_pred.dtype)
@@ -20,6 +20,7 @@
20
20
 
21
21
  from enum import IntEnum
22
22
  import numpy as np
23
+ from scipy.signal import welch
23
24
  from typing import Tuple, Optional, List, Callable, Union
24
25
 
25
26
  from prpy.constants import SECONDS_PER_MINUTE, MILLIS_PER_SECOND
@@ -37,6 +38,14 @@ HR_MIN = 40 # 1/min
37
38
  HR_MAX = 240 # 1/min
38
39
  HRV_SDNN_MIN = 1 # ms
39
40
  HRV_SDNN_MAX = 200 # ms
41
+ HRV_RMSSD_MIN = 1 # ms
42
+ HRV_RMSSD_MAX = 200 # ms
43
+ HRV_LF_MIN = 0 # ms^2
44
+ HRV_LF_MAX = 5000 # ms^2
45
+ HRV_HF_MIN = 0 # ms^2
46
+ HRV_HF_MAX = 5000 # ms^2
47
+ HRV_LF_HF_MIN = 0 # unitless
48
+ HRV_LF_HF_MAX = 10 # unitless
40
49
  PTT_MIN = 100 # ms
41
50
  PTT_MAX = 400 # ms
42
51
  IDX_MIN = 0. # unitless
@@ -50,8 +59,12 @@ SPO2_MAX = 100 # %
50
59
 
51
60
  CALC_HR_MIN_T = 5 # seconds
52
61
  CALC_HR_MAX_T = 10 # seconds
53
- CALC_HRV_SDNN_MIN_T = 10 # seconds
62
+ CALC_HRV_SDNN_MIN_T = 20 # seconds
54
63
  CALC_HRV_SDNN_MAX_T = 60 # seconds
64
+ CALC_HRV_RMSSD_MIN_T = 20 # seconds
65
+ CALC_HRV_RMSSD_MAX_T = 60 # seconds
66
+ CALC_HRV_LF_HF_MIN_T = 55 # seconds
67
+ CALC_HRV_LF_HF_MAX_T = 60 # seconds
55
68
  CALC_RR_MIN_T = 10 # seconds
56
69
  CALC_RR_MAX_T = 30 # seconds
57
70
 
@@ -70,6 +83,14 @@ class EWindowUnit(IntEnum):
70
83
  DETECTIONS = 0
71
84
  SECONDS = 1
72
85
 
86
+ class HRVMetric(IntEnum):
87
+ """Metric for heart rate variability."""
88
+ SDNN = 0
89
+ RMSSD = 1
90
+ LF = 2
91
+ HF = 3
92
+ LF_HF = 4
93
+
73
94
  def estimate_rate_from_signal(
74
95
  signal: np.ndarray,
75
96
  f_s: float,
@@ -523,8 +544,56 @@ def estimate_rr_from_signal(
523
544
  **kw
524
545
  )
525
546
 
526
- def estimate_hrv_sdnn_from_signal(
547
+ def _get_hrv_function(
548
+ metric: HRVMetric,
549
+ f_s: float,
550
+ correct_quant_error: bool = False
551
+ ) -> Callable:
552
+ var_e = 1. / (12 * f_s**2) if correct_quant_error else 0.0
553
+ # SDNN implementation
554
+ def _sdnn_core(diffs: np.ndarray) -> float:
555
+ sd = np.sqrt(np.nanvar(diffs) - var_e) * MILLIS_PER_SECOND
556
+ sd = np.clip(sd, HRV_SDNN_MIN, HRV_SDNN_MAX)
557
+ return sd
558
+ # RMSSD implementation
559
+ def _rmssd_core(diffs: np.ndarray) -> float:
560
+ deltas = np.diff(diffs)
561
+ rmssd = np.sqrt(np.nanmean(deltas**2) - 2*var_e) * MILLIS_PER_SECOND
562
+ rmssd = np.clip(rmssd, HRV_RMSSD_MIN, HRV_RMSSD_MAX)
563
+ return rmssd
564
+ # LF/HF implementation
565
+ def _lf_hf_core(diffs: np.ndarray) -> float:
566
+ t = np.cumsum(diffs, dtype=np.float64)
567
+ t_u = np.arange(0, t[-1], 1 / f_s)
568
+ rr_u = np.interp(t_u, t, diffs)
569
+ freqs, psd = welch(rr_u - rr_u.mean(), fs=f_s, nperseg=256, detrend='linear')
570
+ lf_mask = (freqs >= 0.04) & (freqs < 0.15)
571
+ lf = np.trapz(psd[lf_mask], freqs[lf_mask])
572
+ lf = np.clip(lf, HRV_LF_MIN, HRV_LF_MAX)
573
+ hf_mask = (freqs >= 0.15) & (freqs <= 0.40)
574
+ hf = np.trapz(psd[hf_mask], freqs[hf_mask])
575
+ hf = np.clip(hf, HRV_HF_MIN, HRV_HF_MAX)
576
+ return lf, hf
577
+
578
+ if metric == HRVMetric.SDNN:
579
+ return _sdnn_core
580
+ elif metric == HRVMetric.RMSSD:
581
+ return _rmssd_core
582
+ elif metric == HRVMetric.LF:
583
+ return lambda diffs: _lf_hf_core(diffs)[0]
584
+ elif metric == HRVMetric.HF:
585
+ return lambda diffs: _lf_hf_core(diffs)[1]
586
+ elif metric == HRVMetric.LF_HF:
587
+ def _lf_hf_ratio(diffs: np.ndarray) -> float:
588
+ lf, hf = _lf_hf_core(diffs)
589
+ if hf == 0 or np.isnan(hf):
590
+ return np.nan
591
+ return np.clip(lf / hf, HRV_LF_HF_MIN, HRV_LF_HF_MAX)
592
+ return _lf_hf_ratio
593
+
594
+ def estimate_hrv_from_signal(
527
595
  signal: np.ndarray,
596
+ metric: HRVMetric,
528
597
  f_s: float,
529
598
  min_window_size: float,
530
599
  max_window_size: float,
@@ -545,6 +614,7 @@ def estimate_hrv_sdnn_from_signal(
545
614
 
546
615
  Args:
547
616
  signal: The raw sensor signal. Shape (n,)
617
+ metric: The `HRVMetric` to use.
548
618
  f_s: Sampling frequency [Hz].
549
619
  f_range: Tuple (min, max) of plausible frequency range [Hz].
550
620
  scope: GLOBAL for scalar hrv or ROLLING for hrv trace shape (n,).
@@ -597,8 +667,9 @@ def estimate_hrv_sdnn_from_signal(
597
667
  sdnn = np.nan if scope is EScope.GLOBAL else np.full(signal.shape, np.nan)
598
668
  return sdnn, 0.
599
669
  # Continue using the detections
600
- sdnn = estimate_hrv_sdnn_from_detection_sequences(
670
+ sdnn = estimate_hrv_from_detection_sequences(
601
671
  seqs=det_idxs,
672
+ metric=metric,
602
673
  f_s=f_s,
603
674
  scope=scope,
604
675
  min_window_size=min_window_size,
@@ -612,8 +683,9 @@ def estimate_hrv_sdnn_from_signal(
612
683
  )
613
684
  return sdnn, sdnn_conf
614
685
 
615
- def estimate_hrv_sdnn_from_detections(
686
+ def estimate_hrv_from_detections(
616
687
  det_idxs: np.ndarray,
688
+ metric: HRVMetric,
617
689
  *,
618
690
  f_s: Optional[float] = None,
619
691
  t: Optional[np.ndarray] = None,
@@ -633,6 +705,7 @@ def estimate_hrv_sdnn_from_detections(
633
705
 
634
706
  Args:
635
707
  det_idxs: The detection indices. Shape (n_dets,)
708
+ metric: The `HRVMetric` to use.
636
709
  f_s: The sampling rate. Required when `t` is not given.
637
710
  t: The timestamps of the original signal. Required for scope.ROLLING. Shape (n,)
638
711
  scope: GLOBAL for scalar rate or ROLLING for rate trace shape (n,).
@@ -650,6 +723,7 @@ def estimate_hrv_sdnn_from_detections(
650
723
  - For Scope.ROLLING: Shape (n,)
651
724
  """
652
725
  if f_s is None: f_s = t.shape[0]/(t[-1]-t[0])
726
+ hrv_fn = _get_hrv_function(metric=metric, f_s=f_s, correct_quant_error=correct_quantization_error)
653
727
  def _rate_from_ts(det_t: np.ndarray) -> float:
654
728
  """Convert a 1-D array of detection times to hrv sdnn [ms]"""
655
729
  diffs = np.diff(det_t)
@@ -658,10 +732,7 @@ def estimate_hrv_sdnn_from_detections(
658
732
  if diffs.size - 1 < min_dets or det_t[-1] - det_t[0] < min_t or np.any(diffs == 0):
659
733
  # Signal not sufficient for estimation
660
734
  return np.nan
661
- var_e = 1./(12*f_s**2) if correct_quantization_error else 0
662
- hrv_sdnn = np.sqrt(np.nanvar(diffs) - var_e) * MILLIS_PER_SECOND
663
- if hrv_sdnn > HRV_SDNN_MAX: return HRV_SDNN_MAX
664
- return hrv_sdnn
735
+ return hrv_fn(diffs)
665
736
  def _rate_from_dets(dets: np.ndarray) -> float:
666
737
  """Convert a 1-D array of detection indices to hrv sdnn [ms]"""
667
738
  dets = np.asarray(dets)
@@ -682,8 +753,9 @@ def estimate_hrv_sdnn_from_detections(
682
753
  pad_val=pad_val
683
754
  )
684
755
 
685
- def estimate_hrv_sdnn_from_detection_sequences(
756
+ def estimate_hrv_from_detection_sequences(
686
757
  seqs: List[np.ndarray],
758
+ metric: HRVMetric,
687
759
  *,
688
760
  f_s: Optional[float] = None,
689
761
  t: Optional[np.ndarray] = None,
@@ -703,6 +775,7 @@ def estimate_hrv_sdnn_from_detection_sequences(
703
775
 
704
776
  Args:
705
777
  seqs: List of np.ndarray sequences of detections.
778
+ metric: The `HRVMetric` to use.
706
779
  f_s: The sampling rate. Required when `t` is not given.
707
780
  t: The timestamps of the original signal. Required for scope.ROLLING. Shape (n,)
708
781
  scope: GLOBAL for scalar rate or ROLLING for rate trace shape (n,).
@@ -720,6 +793,7 @@ def estimate_hrv_sdnn_from_detection_sequences(
720
793
  - For Scope.ROLLING: Shape (n,)
721
794
  """
722
795
  if f_s is None: f_s = t.shape[0]/(t[-1]-t[0])
796
+ hrv_fn = _get_hrv_function(metric=metric, f_s=f_s, correct_quant_error=correct_quantization_error)
723
797
  def _rate_from_ts(det_t: np.ndarray) -> float:
724
798
  """Convert a 1-D array of detection times to hrv sdnn [ms]"""
725
799
  diffs = np.diff(det_t)
@@ -728,10 +802,7 @@ def estimate_hrv_sdnn_from_detection_sequences(
728
802
  if diffs.size + 1 < min_dets or det_t[-1] - det_t[0] < min_t or np.any(diffs == 0):
729
803
  # Signal not sufficient for estimation
730
804
  return np.nan
731
- var_e = 1./(12*f_s**2) if correct_quantization_error else 0
732
- hrv_sdnn = np.sqrt(np.nanvar(diffs) - var_e) * MILLIS_PER_SECOND
733
- if hrv_sdnn > HRV_SDNN_MAX: return HRV_SDNN_MAX
734
- return hrv_sdnn
805
+ return hrv_fn(diffs)
735
806
  def _rate_from_dets(dets: np.ndarray) -> float:
736
807
  """Convert a 1-D array of detection indices to hrv sdnn [ms]"""
737
808
  dets = np.asarray(dets)
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: prpy
3
- Version: 0.3.2
3
+ Version: 0.3.4
4
4
  Summary: Collection of Python utils for signal, image, and video processing
5
5
  Author-email: Philipp Rouast <philipp@rouast.com>
6
6
  License: MIT License
@@ -1,4 +1,4 @@
1
- # Copyright (c) 2024 Philipp Rouast
1
+ # Copyright (c) 2025 Philipp Rouast
2
2
  #
3
3
  # Permission is hereby granted, free of charge, to any person obtaining a copy
4
4
  # of this software and associated documentation files (the "Software"), to deal
@@ -24,13 +24,31 @@ sys.path.append('../prpy')
24
24
  import numpy as np
25
25
  import pytest
26
26
 
27
- from prpy.numpy.metric import mag2db, mae, mse, rmse, cor, snr
27
+ from prpy.numpy.metric import _mag2db, mae, mse, rmse, cor, snr
28
+
29
+ def _make_pair(shape, seed=0, nan_ratio=0.25):
30
+ rng = np.random.default_rng(seed)
31
+ y_true = rng.normal(size=shape)
32
+ y_pred = rng.normal(size=shape)
33
+ nan_mask = rng.random(size=shape) < nan_ratio
34
+ y_true[nan_mask] = np.nan
35
+ y_pred[nan_mask & (rng.random(size=shape) < 0.5)] = np.nan
36
+ return y_true, y_pred
37
+
38
+ nanmean_abs = lambda a, b, ax: np.nanmean(np.abs(a - b), axis=ax)
39
+ nanmean_sq = lambda a, b, ax: np.nanmean((a - b) ** 2, axis=ax)
40
+
41
+ def nan_corr_1d(a, b):
42
+ m = np.isfinite(a) & np.isfinite(b)
43
+ if m.sum() < 2: return np.nan
44
+ a, b = a[m] - a[m].mean(), b[m] - b[m].mean()
45
+ return (a * b).mean() / (a.std() * b.std())
28
46
 
29
47
  @pytest.mark.parametrize("shape", [(3,), (2, 3), (2, 3, 5)])
30
48
  def test_mag2db(shape):
31
49
  x = np.zeros(shape=shape)
32
50
  x_copy = x.copy()
33
- out = mag2db(x)
51
+ out = _mag2db(x)
34
52
  assert out.shape == shape
35
53
  np.testing.assert_equal(x, x_copy)
36
54
 
@@ -45,14 +63,14 @@ def test_mae(shape):
45
63
  np.testing.assert_equal(y_true, y_true_copy)
46
64
  np.testing.assert_equal(y_pred, y_pred_copy)
47
65
 
48
- @pytest.mark.parametrize("shape", [(3,), (2, 3), (2, 3, 5)])
49
- def test_mse(shape):
50
- y_true = np.zeros(shape=shape)
51
- y_pred = np.zeros(shape=shape)
52
- y_true_copy = y_true.copy()
53
- y_pred_copy = y_pred.copy()
54
- out = mse(y_true=y_true, y_pred=y_pred, axis=-1)
55
- assert out.shape == shape[:-1]
66
+ @pytest.mark.parametrize("shape", [(4,), (2, 4), (2, 3, 4)])
67
+ def test_mae_ignore_nan(shape):
68
+ y_true, y_pred = _make_pair(shape)
69
+ y_true_copy, y_pred_copy = y_true.copy(), y_pred.copy()
70
+ legacy = mae(y_true, y_pred, axis=-1)
71
+ assert np.isnan(legacy).any()
72
+ ref = nanmean_abs(y_true, y_pred, -1)
73
+ np.testing.assert_allclose(mae(y_true, y_pred, ignore_nan=True), ref, equal_nan=True)
56
74
  np.testing.assert_equal(y_true, y_true_copy)
57
75
  np.testing.assert_equal(y_pred, y_pred_copy)
58
76
 
@@ -67,6 +85,13 @@ def test_rmse(shape):
67
85
  np.testing.assert_equal(y_true, y_true_copy)
68
86
  np.testing.assert_equal(y_pred, y_pred_copy)
69
87
 
88
+ @pytest.mark.parametrize("shape", [(4,), (2,4), (2,3,4)])
89
+ def test_mse_rmse_ignore_nan(shape):
90
+ y_t, y_p = _make_pair(shape, seed=1)
91
+ ref_mse = nanmean_sq(y_t, y_p, -1)
92
+ np.testing.assert_allclose(mse(y_t, y_p, ignore_nan=True), ref_mse, equal_nan=True)
93
+ np.testing.assert_allclose(rmse(y_t, y_p, ignore_nan=True), np.sqrt(ref_mse), equal_nan=True)
94
+
70
95
  @pytest.mark.parametrize("shape", [(3,), (2, 3), (2, 3, 5)])
71
96
  def test_cor(shape):
72
97
  y_true = np.random.uniform(size=shape)
@@ -80,6 +105,14 @@ def test_cor(shape):
80
105
  np.testing.assert_equal(y_true, y_true_copy)
81
106
  np.testing.assert_equal(y_pred, y_pred_copy)
82
107
 
108
+ @pytest.mark.parametrize("shape", [(5,), (3,5), (2,3,5)])
109
+ def test_cor_ignore_nan(shape):
110
+ y_t, y_p = _make_pair(shape, seed=2)
111
+ ref = np.array([nan_corr_1d(a, b)
112
+ for a, b in zip(y_t.reshape(-1, shape[-1]),
113
+ y_p.reshape(-1, shape[-1]))]).reshape(shape[:-1])
114
+ np.testing.assert_allclose(cor(y_t, y_p, ignore_nan=True), ref, equal_nan=True)
115
+
83
116
  @pytest.mark.parametrize("shape", [(6,), (10, 6)])
84
117
  @pytest.mark.parametrize("f_s", [5., 5])
85
118
  @pytest.mark.parametrize("tol", [.2, 1])
@@ -24,13 +24,13 @@ sys.path.append('../prpy')
24
24
  from prpy.constants import SECONDS_PER_MINUTE
25
25
  from prpy.numpy.filters import moving_average, detrend, detrend_frequency_response
26
26
  from prpy.numpy.freq import estimate_freq
27
- from prpy.numpy.physio import EMethod, EScope, EWindowUnit, HR_MIN, HR_MAX
27
+ from prpy.numpy.physio import EMethod, EScope, EWindowUnit, HR_MIN, HR_MAX, HRVMetric
28
28
  from prpy.numpy.physio import estimate_rate_from_signal
29
29
  from prpy.numpy.physio import estimate_rate_from_detections
30
30
  from prpy.numpy.physio import estimate_rate_from_detection_sequences
31
- from prpy.numpy.physio import estimate_hrv_sdnn_from_signal
32
- from prpy.numpy.physio import estimate_hrv_sdnn_from_detections
33
- from prpy.numpy.physio import estimate_hrv_sdnn_from_detection_sequences
31
+ from prpy.numpy.physio import estimate_hrv_from_signal
32
+ from prpy.numpy.physio import estimate_hrv_from_detections
33
+ from prpy.numpy.physio import estimate_hrv_from_detection_sequences
34
34
  from prpy.numpy.physio import moving_average_size_for_hr_response, moving_average_size_for_rr_response
35
35
  from prpy.numpy.physio import detrend_lambda_for_hr_response, detrend_lambda_for_rr_response
36
36
 
@@ -296,15 +296,16 @@ def test_estimate_rate_from_detection_sequences_rolling_dynamic_seconds():
296
296
 
297
297
  def test_estimate_hrv_sdnn_from_detections_global():
298
298
  det_idxs = np.asarray([202, 392, 601, 799, 1201, 1403, 1610, 1839])
299
- actual = estimate_hrv_sdnn_from_detections(det_idxs, f_s=30, interp_skipped=True, min_dets=5, min_t=1.)
299
+ actual = estimate_hrv_from_detections(det_idxs, metric=HRVMetric.SDNN, f_s=30, interp_skipped=True, min_dets=5, min_t=1.)
300
300
  np.testing.assert_allclose(200, actual, atol=0.1)
301
301
 
302
302
  @pytest.mark.parametrize("correct_quantization_error", [False, True])
303
303
  def test_estimate_hrv_sdnn_from_detection_sequences_global(correct_quantization_error):
304
304
  idxs_list = [[202, 392, 612, 799], [1201, 1403, 1610, 1839]]
305
305
  t = np.linspace(0, 8, 2000)
306
- out = estimate_hrv_sdnn_from_detection_sequences(seqs=idxs_list,
306
+ out = estimate_hrv_from_detection_sequences(seqs=idxs_list,
307
307
  t=t,
308
+ metric=HRVMetric.SDNN,
308
309
  correct_quantization_error=correct_quantization_error,
309
310
  scope=EScope.GLOBAL,
310
311
  min_dets=2,
@@ -318,8 +319,9 @@ def test_estimate_hrv_sdnn_from_detection_sequences_global(correct_quantization_
318
319
  def test_estimate_hrv_sdnn_from_detection_sequences_rolling(correct_quantization_error):
319
320
  idxs_list = [[202, 392, 612, 799], [1201, 1403, 1610, 1839]]
320
321
  t = np.linspace(0, 8, 2000)
321
- out = estimate_hrv_sdnn_from_detection_sequences(seqs=idxs_list,
322
+ out = estimate_hrv_from_detection_sequences(seqs=idxs_list,
322
323
  t=t,
324
+ metric=HRVMetric.SDNN,
323
325
  correct_quantization_error=correct_quantization_error,
324
326
  min_window_size=2,
325
327
  max_window_size=4,
@@ -431,8 +433,9 @@ def test_estimate_hrv_sdnn_from_signal_with_confidence(scenario):
431
433
  signal, conf, f_s, conf_threshold, expected, exp_conf = scenario
432
434
  signal = np.asarray(signal)
433
435
  conf = np.asarray(conf)
434
- actual, conf = estimate_hrv_sdnn_from_signal(
436
+ actual, conf = estimate_hrv_from_signal(
435
437
  signal=signal,
438
+ metric=HRVMetric.SDNN,
436
439
  f_s=f_s,
437
440
  min_window_size=27,
438
441
  max_window_size=27,
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes