prpy 0.2.2__py3-none-any.whl

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.
prpy/numpy/face.py ADDED
@@ -0,0 +1,213 @@
1
+ # Copyright (c) 2024 Philipp Rouast
2
+ #
3
+ # Permission is hereby granted, free of charge, to any person obtaining a copy
4
+ # of this software and associated documentation files (the "Software"), to deal
5
+ # in the Software without restriction, including without limitation the rights
6
+ # to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
7
+ # copies of the Software, and to permit persons to whom the Software is
8
+ # furnished to do so, subject to the following conditions:
9
+ #
10
+ # The above copyright notice and this permission notice shall be included in all
11
+ # copies or substantial portions of the Software.
12
+ #
13
+ # THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
14
+ # IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
15
+ # FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
16
+ # AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
17
+ # LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
18
+ # OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
19
+ # SOFTWARE.
20
+
21
+ import numpy as np
22
+ from prpy.numpy.image import crop_slice_resize
23
+ from typing import Tuple, Union
24
+
25
+ def _get_roi_from_det(
26
+ det: tuple,
27
+ rel_change: tuple,
28
+ clip_dims: Union[tuple, None] = None
29
+ ) -> Tuple[int, int, int, int]:
30
+ """Convert face detection to roi by relative add/reduce.
31
+
32
+ Args:
33
+ det: The face detection [0, H/W] in format (x0, y0, x1, y1).
34
+ rel_change: The relative change to make in format (left, top, right, bottom).
35
+ clip_dims: tuple (frame_w, frame_h) to clip the result to (optional).
36
+ Returns:
37
+ out: The roi [0, H/W] in format (x0, y0, x1, y1)
38
+ """
39
+ assert isinstance(det, tuple) and len(det) == 4 and all(isinstance(i, int) for i in det)
40
+ assert det[2] > det[0]
41
+ assert det[3] > det[1]
42
+ assert isinstance(rel_change, tuple) and len(rel_change) == 4 and all(isinstance(i, float) for i in rel_change)
43
+ assert clip_dims is None or (isinstance(clip_dims, tuple) and len(clip_dims) == 2 and all(isinstance(i, int) for i in clip_dims))
44
+ def _clip_dims(val, min_dim, max_dim):
45
+ return min(max(val, min_dim), max_dim)
46
+ det_w = det[2]-det[0]
47
+ det_h = det[3]-det[1]
48
+ rel_ch_l, rel_ch_t, rel_ch_r, rel_ch_b = rel_change
49
+ abs_ch_l = int(rel_ch_l * det_w)
50
+ abs_ch_t = int(rel_ch_t * det_h)
51
+ abs_ch_r = int(rel_ch_r * det_w)
52
+ abs_ch_b = int(rel_ch_b * det_h)
53
+ if clip_dims is not None:
54
+ return (_clip_dims(det[0] - abs_ch_l, 0, clip_dims[0]),
55
+ _clip_dims(det[1] - abs_ch_t, 0, clip_dims[1]),
56
+ _clip_dims(det[2] + abs_ch_r, 0, clip_dims[0]),
57
+ _clip_dims(det[3] + abs_ch_b, 0, clip_dims[1]))
58
+ else:
59
+ return (det[0]-abs_ch_l, det[1]-abs_ch_t, det[2]+abs_ch_r, det[3]+abs_ch_b)
60
+
61
+ def get_face_roi_from_det(det: tuple) -> tuple:
62
+ """Convert face detection into face roi.
63
+ Reduces width to 60% and height to 80%.
64
+
65
+ Args:
66
+ det: The face detection [0, H/W] in form (x0, y0, x1, y1)
67
+ Returns:
68
+ out: The roi [0, H/W] in form (x0, y0, x1, y1)
69
+ """
70
+ return _get_roi_from_det(det=det, rel_change=(-0.2, -0.1, -0.2, -0.1))
71
+
72
+ def get_forehead_roi_from_det(det: tuple) -> tuple:
73
+ """Convert face detection into forehead roi.
74
+ Reduces det to forehead as 35% to 65% of width, and 15% to 25% of height.
75
+
76
+ Args:
77
+ det: The face detection [0, H/W] in form (x0, y0, x1, y1)
78
+ Returns:
79
+ out: The roi [0, H/W] in form (x0, y0, x1, y1)
80
+ """
81
+ return _get_roi_from_det(det=det, rel_change=(-0.35, -0.15, -0.35, -0.75))
82
+
83
+ def get_upper_body_roi_from_det(
84
+ det: tuple,
85
+ clip_dims: tuple,
86
+ cropped: bool = False,
87
+ v: int = 1
88
+ ) -> tuple:
89
+ """Convert face detection into upper body roi and clip to frame constraints.
90
+
91
+ Args:
92
+ det: The face detection [0, H/W] in form (x0, y0, x1, y1)
93
+ clip_dims: constraints (frame_w, frame_h) to clip the result to
94
+ cropped: Create cropped variant?
95
+ v: Version of ROI definition (0, 1, 2, or 3)
96
+ Returns:
97
+ out: The roi [0, H/W] in form (x0, y0, x1, y1)
98
+ """
99
+ assert isinstance(cropped, bool)
100
+ assert isinstance(v, int)
101
+ if v == 0:
102
+ # V0: (.25, .3, .25, .5) -> (.175, .27, .175, .45)
103
+ if not cropped:
104
+ return _get_roi_from_det(
105
+ det=det, rel_change=(.25, .3, .25, .5), clip_dims=clip_dims)
106
+ else:
107
+ return _get_roi_from_det(
108
+ det=det, rel_change=(.175, .27, .175, .45), clip_dims=clip_dims)
109
+ elif v == 1:
110
+ # V1: (.25, .2, .25, .4) -> (.175, .15, .175, .3)
111
+ if not cropped:
112
+ return _get_roi_from_det(
113
+ det=det, rel_change=(.25, .2, .25, .4), clip_dims=clip_dims)
114
+ else:
115
+ return _get_roi_from_det(
116
+ det=det, rel_change=(.175, .15, .175, .3), clip_dims=clip_dims)
117
+ elif v == 2:
118
+ # V2: (.25, .1, .25, .5) -> (.175, .075, .175, .375)
119
+ if not cropped:
120
+ return _get_roi_from_det(
121
+ det=det, rel_change=(.25, .1, .25, .5), clip_dims=clip_dims)
122
+ else:
123
+ return _get_roi_from_det(
124
+ det=det, rel_change=(.175, .075, .175, .375), clip_dims=clip_dims)
125
+ elif v == 3:
126
+ # V3: (.2, .3, .2, .45) -> (.15, .25, .15, .35)
127
+ if not cropped:
128
+ return _get_roi_from_det(
129
+ det=det, rel_change=(.2, .3, .2, .45), clip_dims=clip_dims)
130
+ else:
131
+ return _get_roi_from_det(
132
+ det=det, rel_change=(.15, .25, .15, .35), clip_dims=clip_dims)
133
+ else:
134
+ raise ValueError("v {} is not defined".format(v))
135
+
136
+ def get_meta_roi_from_det(
137
+ det: tuple,
138
+ clip_dims: tuple
139
+ ) -> tuple:
140
+ """Convert face detection into meta roi and clip to frame constraints.
141
+
142
+ Args:
143
+ det: The face detection [0, H/W] in form (x0, y0, x1, y1)
144
+ clip_dims: constraints (frame_w, frame_h) to clip the result to
145
+ Returns:
146
+ out: The roi [0, H/W] in form (x0, y0, x1, y1)
147
+ """
148
+ return _get_roi_from_det(
149
+ det=det, rel_change=(.2, .2, .2, .2), clip_dims=clip_dims)
150
+
151
+ def get_roi_from_det(
152
+ det: tuple,
153
+ roi_method: Union[str, None],
154
+ clip_dims: Union[tuple, None] = None
155
+ ) -> tuple:
156
+ """Convert face detection into specified roi.
157
+
158
+ Args:
159
+ det: The face detection [0, H/W] in form (x0, y0, x1, y1)
160
+ roi_method: Which roi method to use. Either 'forehead', 'face',
161
+ 'upper_body', 'upper_body_cropped', 'meta', None (directly use det)
162
+ clip_dims: Constraints (frame_w, frame_h) to clip the result to (optional).
163
+ Returns:
164
+ out: The roi [0, H/W] in form (x0, y0, x1, y1)
165
+ """
166
+ assert roi_method is None or isinstance(roi_method, str)
167
+ if roi_method == 'face':
168
+ return get_face_roi_from_det(det)
169
+ elif roi_method == 'forehead':
170
+ return get_forehead_roi_from_det(det)
171
+ elif roi_method == 'upper_body':
172
+ assert clip_dims is not None
173
+ return get_upper_body_roi_from_det(det, clip_dims=clip_dims, cropped=False)
174
+ elif roi_method == 'upper_body_cropped':
175
+ assert clip_dims is not None
176
+ return get_upper_body_roi_from_det(det, clip_dims=clip_dims, cropped=True)
177
+ elif roi_method == 'meta':
178
+ assert clip_dims is not None
179
+ return get_meta_roi_from_det(det, clip_dims=clip_dims)
180
+ elif roi_method is None or roi_method == 'det':
181
+ return det
182
+ else:
183
+ raise ValueError("roi method {} is not supported".format(roi_method))
184
+
185
+ def crop_resize_from_det(
186
+ video: np.ndarray,
187
+ det: tuple,
188
+ size: tuple,
189
+ roi_method: str,
190
+ library: str,
191
+ scale_algorithm: str
192
+ ) -> np.ndarray:
193
+ """Crop and resize a video according to a single face detection.
194
+ Resize to specified size with specified method.
195
+
196
+ Args:
197
+ video: The video. Shape (n_frames, h, w, c)
198
+ det: The face detection in form (x_0, y_0, x_1, y_1)
199
+ size: The target size for resize - (h, w)
200
+ roi_method: Which roi method to use. Either 'forehead', 'face',
201
+ 'upper_body', 'upper_body_cropped', 'meta', None (directly use det)
202
+ library: The library used for resize (PIL, cv2, or tf - returns tf.Tensor)
203
+ scale_algorithm: The algorithm used for scaling. Supports: bicubic,
204
+ bilinear, area (not for PIL!), lanczos
205
+ Returns:
206
+ result: Cropped and resized video. Shape [n_frames, size[0], size[1], c]
207
+ """
208
+ assert isinstance(video, np.ndarray) and len(video.shape) == 4
209
+ _, height, width, _ = video.shape
210
+ roi = get_roi_from_det(det, roi_method=roi_method, clip_dims=(width, height))
211
+ return crop_slice_resize(
212
+ inputs=video, target_size=size, roi=roi, library=library,
213
+ preserve_aspect_ratio=False, scale_algorithm=scale_algorithm)
prpy/numpy/image.py ADDED
@@ -0,0 +1,141 @@
1
+ # Copyright (c) 2024 Philipp Rouast
2
+ #
3
+ # Permission is hereby granted, free of charge, to any person obtaining a copy
4
+ # of this software and associated documentation files (the "Software"), to deal
5
+ # in the Software without restriction, including without limitation the rights
6
+ # to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
7
+ # copies of the Software, and to permit persons to whom the Software is
8
+ # furnished to do so, subject to the following conditions:
9
+ #
10
+ # The above copyright notice and this permission notice shall be included in all
11
+ # copies or substantial portions of the Software.
12
+ #
13
+ # THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
14
+ # IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
15
+ # FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
16
+ # AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
17
+ # LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
18
+ # OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
19
+ # SOFTWARE.
20
+
21
+ import numpy as np
22
+ from typing import Union
23
+
24
+ def crop_slice_resize(
25
+ inputs: np.ndarray,
26
+ target_size: Union[int, tuple, list],
27
+ roi: Union[tuple, list, None] = None,
28
+ target_idxs: Union[tuple, list, np.ndarray, None] = None,
29
+ preserve_aspect_ratio: bool = False,
30
+ keepdims: bool = True,
31
+ library: str = 'PIL',
32
+ scale_algorithm: str = 'bicubic'
33
+ ) -> np.ndarray:
34
+ """Crop, slice, and resize image(s) with all same settings.
35
+
36
+ Args:
37
+ inputs: The inputs as uint8 shape (h, w, 3) or (n_frames, h, w, 3)
38
+ target_size: The target size; scalar or (H, W) if preserve_aspect_ratio=False
39
+ roi: The region of interest in format (x0, y0, x1, y1). Use None to keep all.
40
+ target_idxs: The frame indices to be used. Use None to keep all.
41
+ preserve_aspect_ratio: Preserve the aspect ratio?
42
+ keepdims: If True, always keep n_frames dim. Otherwise, may drop n_frames dim.
43
+ library: The library to use. `cv2` or `PIL` (return np ndarray), or `tf` (returns tf Tensor)
44
+ scale_algorithm: The algorithm used for scaling.
45
+ Supports: bicubic, bilinear, area (not for PIL!), lanczos. Default: bicubic
46
+ Returns:
47
+ result: The processed frame(s) as float32 shape (h, w, 3) or (n_frames, h, w, 3)
48
+ """
49
+ assert isinstance(inputs, np.ndarray) and (len(inputs.shape) == 3 or len(inputs.shape) == 4)
50
+ assert isinstance(target_size, int) or (isinstance(target_size, (tuple, list)) and len(target_size) == 2 and all(isinstance(i, int) for i in target_size))
51
+ assert roi is None or (isinstance(roi, (tuple, list)) and len(roi) == 4 and all(isinstance(i, int) for i in roi) and roi[2] > roi[0] and roi[3] > roi[1])
52
+ assert target_idxs is None or isinstance(target_idxs, np.ndarray) or (isinstance(target_idxs, (tuple, list)) and all(isinstance(i, int) for i in target_idxs))
53
+ assert isinstance(preserve_aspect_ratio, bool)
54
+ assert isinstance(keepdims, bool)
55
+ assert isinstance(library, str)
56
+ assert isinstance(scale_algorithm, str)
57
+ unpack_target_size = lambda x: (x[0], x[1]) if isinstance(x, (list, tuple)) else (x, x)
58
+ target_height, target_width = unpack_target_size(target_size)
59
+ inputs_shape = inputs.shape
60
+ # Add temporal dim if necessary
61
+ if len(inputs_shape) == 3: inputs = inputs[np.newaxis,:,:,:]
62
+ # Apply target_idxs and roi
63
+ inputs = inputs[(target_idxs if target_idxs is not None else slice(None)),
64
+ (slice(roi[1], roi[3]) if isinstance(roi, (tuple, list)) else slice(None)),
65
+ (slice(roi[0], roi[2]) if isinstance(roi, (tuple, list)) else slice(None))]
66
+ in_shape = inputs.shape
67
+ # Compute out size
68
+ def _out_size(in_shape, target_height, target_width, preserve_aspect_ratio):
69
+ _, height, width, _ = in_shape
70
+ if preserve_aspect_ratio:
71
+ # Determine critical side
72
+ inputs_r = float(width)/height
73
+ target_r = float(target_width)/target_height
74
+ if inputs_r < target_r:
75
+ # Height is critical side
76
+ out_size = (target_height, int(target_height*inputs_r))
77
+ else:
78
+ # Height is critical side
79
+ out_size = (int(target_width/inputs_r), target_width)
80
+ else:
81
+ out_size = (target_height, target_width)
82
+ return out_size
83
+ out_size = _out_size(in_shape, target_height, target_width, preserve_aspect_ratio)
84
+ # Distinguish between different cases
85
+ if out_size == (in_shape[1], in_shape[2]):
86
+ # No resizing necessary
87
+ if library == 'tf':
88
+ import tensorflow as tf
89
+ out = tf.convert_to_tensor(inputs, dtype=tf.float32)
90
+ else:
91
+ out = inputs.astype(np.float32)
92
+ else:
93
+ # Resize to out_size
94
+ if library == 'tf':
95
+ import tensorflow as tf
96
+ # https://www.tensorflow.org/api_docs/python/tf/image/ResizeMethod
97
+ mapping = {"bicubic": "bicubic", "bilinear": "bilinear", "lanczos": "lanczos3", "area": "area"}
98
+ try:
99
+ library_algorithm = mapping[scale_algorithm]
100
+ except KeyError:
101
+ raise ValueError("Scaling algorithm {} is not supported by {}".format(scale_algorithm, library))
102
+ out = tf.image.resize(
103
+ images=inputs, size=(target_height, target_width),
104
+ preserve_aspect_ratio=preserve_aspect_ratio,
105
+ method=library_algorithm, antialias=False)
106
+ elif library == 'PIL':
107
+ from PIL import Image
108
+ # https://pillow.readthedocs.io/en/stable/releasenotes/2.7.0.html#image-resizing-filters
109
+ mapping = {"bicubic": Image.BICUBIC, "bilinear": Image.BILINEAR, "lanczos": Image.LANCZOS}
110
+ try:
111
+ library_algorithm = mapping[scale_algorithm]
112
+ except KeyError:
113
+ raise ValueError("Scaling algorithm {} is not supported by {}".format(scale_algorithm, library))
114
+ # PIL requires (width, height)
115
+ out_size = (out_size[1], out_size[0])
116
+ out = np.asarray([
117
+ np.asarray(Image.fromarray(f).resize(out_size, resample=library_algorithm)) for f in inputs])
118
+ elif library == 'cv2':
119
+ import cv2
120
+ # https://docs.opencv.org/3.4/da/d54/group__imgproc__transform.html
121
+ mapping = {"bicubic": cv2.INTER_CUBIC, "bilinear": cv2.INTER_LINEAR, "lanczos": cv2.INTER_LANCZOS4, "area": cv2.INTER_AREA}
122
+ try:
123
+ library_algorithm = mapping[scale_algorithm]
124
+ except KeyError:
125
+ raise ValueError("Scaling algorithm {} is not supported by {}".format(scale_algorithm, library))
126
+ # cv2 requires (width, height)
127
+ out_size = (out_size[1], out_size[0])
128
+ out = np.asarray([
129
+ cv2.resize(src=f, dsize=out_size, interpolation=library_algorithm) for f in inputs])
130
+ else:
131
+ raise ValueError("Library {} not supported".format(library))
132
+ if keepdims and len(out.shape) == 3 and len(inputs_shape) == 4:
133
+ # Add temporal dim back - might have been lost when slicing
134
+ newaxis = tf.newaxis if library == 'tf' else np.newaxis
135
+ out = out[newaxis,:,:,:]
136
+ elif not keepdims and len(out.shape) == 4:
137
+ # Remove temporal dim if necessary
138
+ squeeze = tf.squeeze if library == 'tf' else np.squeeze
139
+ if out.shape[0] == 1:
140
+ out = squeeze(out, axis=0)
141
+ return out
prpy/numpy/metric.py ADDED
@@ -0,0 +1,179 @@
1
+ # Copyright (c) 2024 Philipp Rouast
2
+ #
3
+ # Permission is hereby granted, free of charge, to any person obtaining a copy
4
+ # of this software and associated documentation files (the "Software"), to deal
5
+ # in the Software without restriction, including without limitation the rights
6
+ # to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
7
+ # copies of the Software, and to permit persons to whom the Software is
8
+ # furnished to do so, subject to the following conditions:
9
+ #
10
+ # The above copyright notice and this permission notice shall be included in all
11
+ # copies or substantial portions of the Software.
12
+ #
13
+ # THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
14
+ # IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
15
+ # FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
16
+ # AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
17
+ # LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
18
+ # OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
19
+ # SOFTWARE.
20
+
21
+ import numpy as np
22
+ from scipy import signal
23
+ from typing import Union
24
+
25
+ def mag2db(mag: Union[np.ndarray, np.float64]) -> np.ndarray:
26
+ """Magnitude to decibels element-wise.
27
+
28
+ Args:
29
+ mag: Magnitude. Arbitrary shape.
30
+ Returns:
31
+ out: Decibels. Same shape as input.
32
+ """
33
+ assert isinstance(mag, (np.ndarray, np.float64))
34
+ return 20. * np.log10(mag)
35
+
36
+ def mae(
37
+ y_true: np.ndarray,
38
+ y_pred: np.ndarray,
39
+ axis: Union[int, None] = -1
40
+ ) -> np.ndarray:
41
+ """Mean absolute error
42
+
43
+ Args:
44
+ y_true: True values. Shape (..., dim_n, dim_axis)
45
+ y_pred: Predicted values. Shape (..., dim_n, dim_axis)
46
+ axis: Axis along which the means are computed
47
+ Returns:
48
+ mae: The mean absolute error. Shape (..., dim_n)
49
+ """
50
+ assert axis is None or isinstance(axis, int)
51
+ y_true = np.asarray(y_true)
52
+ y_pred = np.asarray(y_pred)
53
+ assert y_true.shape == y_pred.shape
54
+ if np.all(np.isfinite(y_true)) and np.all(np.isfinite(y_pred)):
55
+ return np.mean(np.abs(y_true - y_pred), axis=axis)
56
+ else:
57
+ return np.full(y_true.shape[:axis], fill_value=np.nan, dtype=y_true.dtype)
58
+
59
+ def mse(
60
+ y_true: np.ndarray,
61
+ y_pred: np.ndarray,
62
+ axis: Union[int, None] = -1
63
+ ) -> np.ndarray:
64
+ """Mean squared error
65
+
66
+ Args:
67
+ y_true: True values. Shape (..., dim_n, dim_axis)
68
+ y_pred: Predicted values. Shape (..., dim_n, dim_axis)
69
+ axis: Axis along which the means are computed
70
+ Returns:
71
+ mse: The mean squared error. Shape (..., dim_n)
72
+ """
73
+ assert axis is None or isinstance(axis, int)
74
+ y_true = np.asarray(y_true)
75
+ y_pred = np.asarray(y_pred)
76
+ assert y_true.shape == y_pred.shape
77
+ if np.all(np.isfinite(y_true)) and np.all(np.isfinite(y_pred)):
78
+ return np.mean(np.square(y_true - y_pred), axis=axis)
79
+ else:
80
+ return np.full(y_true.shape[:axis], fill_value=np.nan, dtype=y_true.dtype)
81
+
82
+ def rmse(
83
+ y_true: np.ndarray,
84
+ y_pred: np.ndarray,
85
+ axis: Union[int, None] = -1
86
+ ) -> np.ndarray:
87
+ """Root mean squared error
88
+
89
+ Args:
90
+ y_true: True values. Shape (..., dim_n, dim_axis)
91
+ y_pred: Predicted values. Shape (..., dim_n, dim_axis)
92
+ axis: Axis along which the means are computed
93
+ Returns:
94
+ rmse: The root mean squared error. Shape (..., dim_n)
95
+ """
96
+ assert axis is None or isinstance(axis, int)
97
+ y_true = np.asarray(y_true)
98
+ y_pred = np.asarray(y_pred)
99
+ assert y_true.shape == y_pred.shape
100
+ if np.all(np.isfinite(y_true)) and np.all(np.isfinite(y_pred)):
101
+ return np.sqrt(np.mean(np.square(y_true - y_pred), axis=axis))
102
+ else:
103
+ return np.full(y_true.shape[:axis], fill_value=np.nan, dtype=y_true.dtype)
104
+
105
+ def cor(
106
+ y_true: np.ndarray,
107
+ y_pred: np.ndarray,
108
+ axis: Union[int, None] = -1
109
+ ) -> np.ndarray:
110
+ """Pearson's correlation coefficient
111
+
112
+ Args:
113
+ y_true: True values. Shape (..., dim_n, dim_axis)
114
+ y_pred: Predicted values. Shape (..., dim_n, dim_axis)
115
+ axis: Axis along which correlations are computed
116
+ Returns:
117
+ cor: The correlation coefficients. Shape (..., dim_n)
118
+ """
119
+ assert axis is None or isinstance(axis, int)
120
+ y_true = np.asarray(y_true)
121
+ y_pred = np.asarray(y_pred)
122
+ assert y_true.shape == y_pred.shape
123
+ if np.all(np.isfinite(y_true)) and np.all(np.isfinite(y_pred)):
124
+ res_true = y_true - np.mean(y_true, axis=axis, keepdims=True)
125
+ res_pred = y_pred - np.mean(y_pred, axis=axis, keepdims=True)
126
+ cov = np.mean(res_true * res_pred, axis=axis)
127
+ var_true = np.mean(res_true**2, axis=axis)
128
+ var_pred = np.mean(res_pred**2, axis=axis)
129
+ sigma_true = np.sqrt(var_true)
130
+ sigma_pred = np.sqrt(var_pred)
131
+ return cov / (sigma_true * sigma_pred)
132
+ else:
133
+ return np.full(y_true.shape[:axis], fill_value=np.nan, dtype=y_true.dtype)
134
+
135
+ def snr(
136
+ f_true: np.ndarray,
137
+ y_pred: np.ndarray,
138
+ f_s: float,
139
+ f_res: float,
140
+ tol: float = .1,
141
+ f_min: float = .5,
142
+ f_max: float = 4.):
143
+ """Signal-to-noise ratio
144
+ Args:
145
+ f_true: The true frequencies. Shape (b,) or ()
146
+ y_pred: Predicted vals. Shape (b, t) or (t,)
147
+ f_s: Sampling frequency
148
+ f_res: Frequency resolution
149
+ tol: Frequency domain tolerance
150
+ f_min: Minimum frequency included in metric calculation
151
+ f_max: Maximum frequency included in metric calculation
152
+ Returns:
153
+ snr: The signal to noise ratio. Shape (b,) or ()
154
+ """
155
+ assert isinstance(f_s, float)
156
+ assert isinstance(f_res, float)
157
+ assert isinstance(tol, float)
158
+ assert isinstance(f_min, float)
159
+ assert isinstance(f_max, float)
160
+ f_true = np.asarray(f_true)
161
+ y_pred = np.asarray(y_pred)
162
+ assert len(y_pred.shape) == 1 or len(y_pred.shape) == 2
163
+ assert f_true.shape == y_pred.shape[:-1]
164
+ if np.all(np.isfinite(y_pred)) and np.all(np.isfinite(f_true)):
165
+ n = f_s // f_res
166
+ f, pxx = signal.periodogram(y_pred, fs=f_s, nfft=n, detrend=False, axis=-1)
167
+ if len(y_pred.shape) == 2:
168
+ f = np.broadcast_to(f[np.newaxis], pxx.shape)
169
+ f_true = f_true[...,np.newaxis]
170
+ gt_mask_1 = (f >= f_true - tol) & (f <= f_true + tol)
171
+ gt_mask_2 = (f >= f_true * 2 - tol) & (f <= f_true * 2 + tol)
172
+ gt_mask = gt_mask_1 | gt_mask_2
173
+ s_power = np.sum(pxx * gt_mask, axis=-1)
174
+ f_mask = (f >= f_min) & (f <= f_max)
175
+ all_power = np.sum(pxx * f_mask, axis=-1)
176
+ snr = mag2db(s_power / (all_power - s_power))
177
+ return np.squeeze(snr)
178
+ else:
179
+ return np.full(y_pred.shape[:-1], fill_value=np.nan, dtype=y_pred.dtype)