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.
@@ -0,0 +1,182 @@
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 window_view(
25
+ x: np.ndarray,
26
+ min_window_size: int,
27
+ max_window_size: int,
28
+ overlap: int,
29
+ pad_mode: str = 'constant',
30
+ const_val: Union[float, int] = np.nan
31
+ ) -> np.ndarray:
32
+ """Create a window view into an n-d array `x` along its first dim.
33
+
34
+ Args:
35
+ x: The n-d array into which we want to create a windowed view
36
+ min_window_size: The minimum window size
37
+ max_window_size: The maximum window size
38
+ overlap: The overlap of the sliding windows
39
+ pad_mode: The pad mode
40
+ const_val: The constant value to be padded with if pad_mode == 'constant'
41
+ Returns:
42
+ y: The n+1-d windowed view into x of shape (n_windows, window_size, ...)
43
+ pad_start: How much padding was applied at the start (scalar)
44
+ pad_end: How much padding was applied at the end (scalar)
45
+ """
46
+ assert isinstance(min_window_size, int) and min_window_size >= 0
47
+ assert isinstance(max_window_size, int) and max_window_size >= min_window_size
48
+ assert isinstance(overlap, int) and overlap < max_window_size, "overlap must be smaller than max_window_size"
49
+ assert isinstance(pad_mode, str)
50
+ assert isinstance(const_val, (float, int))
51
+ x = np.asarray(x)
52
+ original_len = x.shape[0]
53
+ step_size = max_window_size - overlap
54
+ # Pad the start using variable window sizes
55
+ pad_start = max_window_size - min_window_size
56
+ # Pad the end if there is a remainder
57
+ remainder = (original_len - min_window_size) % step_size
58
+ pad_end = 0 if remainder == 0 else step_size - remainder
59
+ pad_width = ((pad_start, pad_end),) + ((0, 0),) * (x.ndim - 1)
60
+ if pad_mode=='constant':
61
+ x = np.pad(x, pad_width, mode=pad_mode, constant_values=(const_val,))
62
+ else:
63
+ x = np.pad(x, pad_width, mode=pad_mode)
64
+ # Calculate the shape of the view for the first dimension
65
+ new_shape = ((pad_start + original_len + pad_end - max_window_size) // step_size + 1, max_window_size)
66
+ # Add shapes for following dimensions
67
+ new_shape += x.shape[1:]
68
+ # Calculate the stride of the view for the first dimension
69
+ new_strides = ((step_size * x.strides[0],) + (x.strides[0],))
70
+ # Add strides for following dimensions
71
+ new_strides += x.strides[1:]
72
+ # Generate the view
73
+ y = np.lib.stride_tricks.as_strided(
74
+ x, shape=new_shape, strides=new_strides)
75
+ # Return
76
+ return y, pad_start, pad_end
77
+
78
+ def reduce_window_view(
79
+ x: np.ndarray,
80
+ overlap: int,
81
+ pad_end: int = 0,
82
+ hanning: bool = False
83
+ ) -> np.ndarray:
84
+ """Reduce an n-d window view by arranging the first dimension as sliding
85
+ windows and then reducing it using the mean.
86
+
87
+ Args:
88
+ x: The n-d window view of shape (n_windows, window_size, ...)
89
+ overlap: The overlap with which the window view was created
90
+ pad_end: How much padding was applied to the end when the window view was created
91
+ hanning: Whether to reduce the window view with hanning windows
92
+ Returns:
93
+ mean: The n-1-d reduced window view [original_len, ...]
94
+ """
95
+ assert isinstance(x, np.ndarray)
96
+ assert isinstance(hanning, bool)
97
+ # Infer number of windows and original length (including extra padding)
98
+ num_windows = x.shape[0]
99
+ window_size = x.shape[1]
100
+ assert isinstance(overlap, int) and overlap >= 0 and overlap < window_size
101
+ assert isinstance(pad_end, int) and pad_end >= 0 and pad_end < window_size
102
+ original_len_with_pad_end = num_windows * window_size - (num_windows-1) * overlap
103
+ # Apply hanning window to taper x
104
+ if hanning:
105
+ x *= np.hanning(window_size)
106
+ # Add padding to extend the matrix to the dimensions of the diagonal matrix
107
+ padding = ((0, 0), (0, original_len_with_pad_end - window_size)) + ((0, 0),) * (x.ndim - 2)
108
+ y = np.pad(x, padding, 'constant')
109
+ # Use as_strided to create a view of the diagonal matrix
110
+ # that aligns the windows temporally as they were generated
111
+ # https://stackoverflow.com/a/60460462/3595278
112
+ y_roll = y[:, [*range(y.shape[1]),*range(y.shape[1]-1)]].copy() #need `copy`
113
+ view_strides = list(y_roll.strides)
114
+ view_strides.insert(1, y_roll.strides[1])
115
+ view_strides = tuple(view_strides)
116
+ view_shape = list(y.shape)
117
+ view_shape.insert(1, original_len_with_pad_end)
118
+ view_shape = tuple(view_shape)
119
+ view = np.lib.stride_tricks.as_strided(y_roll, view_shape, view_strides)
120
+ step_size = window_size - overlap
121
+ m = np.asarray([step_size * i for i in range(num_windows)])
122
+ view = view[np.arange(y.shape[0]), (original_len_with_pad_end-m)%original_len_with_pad_end]
123
+ # Merge the windows by taking the mean across the time dimension
124
+ mean = np.true_divide(
125
+ view.sum(0), np.maximum((view != 0).sum(0), 1))
126
+ # Trim result by pad_end if necessary
127
+ if pad_end > 0: mean = mean[:-pad_end]
128
+ return mean
129
+
130
+ def resolve_1d_window_view(
131
+ x: np.ndarray,
132
+ window_size: int,
133
+ overlap: int,
134
+ pad_end: int,
135
+ fill_method: str
136
+ ) -> np.ndarray:
137
+ """Resolve an 1-d window view by extending it to the expected shape.
138
+
139
+ - This is useful if processing on each window created a scalar value.
140
+
141
+ Args:
142
+ x: The 1-d window view to be resolved
143
+ window_size: The window size used to create the view
144
+ overlap: The overlap used to create the view
145
+ pad_end: How much padding was applied to the end
146
+ fill_method: Method for filling/padding the result
147
+ Returns:
148
+ vals: The 1-d resolved data
149
+ """
150
+ assert isinstance(x, np.ndarray) and len(x.shape) == 1
151
+ assert isinstance(window_size, int) and window_size > 0
152
+ assert isinstance(overlap, int) and overlap >= 0 and overlap < window_size
153
+ assert isinstance(pad_end, int) and pad_end >= 0 and pad_end < window_size
154
+ assert isinstance(fill_method, str)
155
+ if overlap == 0:
156
+ # If overlap is zero, we simply need to repeat each value to match window_size
157
+ vals = np.repeat(x, window_size)
158
+ if pad_end > 0:
159
+ # Trim end if it has been padded
160
+ vals = vals[:-pad_end]
161
+ elif overlap == window_size - 1:
162
+ # If overlap is one less than the window size, then values are mostly
163
+ # already ok. We just have to fill the start.
164
+ if fill_method == 'zero':
165
+ fill = 0.0
166
+ elif fill_method == 'mean':
167
+ fill = np.mean(x)
168
+ elif fill_method == 'start':
169
+ fill = x[0]
170
+ else:
171
+ raise ValueError("fill_method {} not supported".format(fill_method))
172
+ vals = np.concatenate([np.repeat(fill, window_size-1), x])
173
+ elif overlap < window_size - 1:
174
+ # For any other overlaps, we will have to build an intermediate 2-d
175
+ # window view of the vals and reduce it via reduce_window_view.
176
+ n = len(x)
177
+ x = np.reshape(np.repeat(x, window_size), (n, window_size))
178
+ vals = reduce_window_view(x, overlap=overlap)
179
+ if pad_end > 0:
180
+ # Trim end if it has been padded
181
+ vals = vals[:-pad_end]
182
+ return vals
@@ -0,0 +1,19 @@
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.
@@ -0,0 +1,203 @@
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
+ import tensorflow as tf
23
+ from typing import Union
24
+
25
+ def _reduction_dims(
26
+ x: tf.Tensor,
27
+ axis: Union[int, tuple, None]
28
+ ) -> Union[int, tuple, tf.Tensor]:
29
+ """Resolve the reduction dims.
30
+
31
+ - If axis is None, returns all dims of x for reduction
32
+
33
+ Args:
34
+ x: The tensor to be reduced
35
+ axis: Dims to be reduced or None
36
+ """
37
+ assert isinstance(x, tf.Tensor)
38
+ assert axis is None or isinstance(axis, int) or isinstance(axis, tuple)
39
+ if axis is not None:
40
+ return axis
41
+ else:
42
+ x_rank = None
43
+ if isinstance(x, tf.Tensor):
44
+ x_rank = x.shape.rank
45
+ # Fast path: avoid creating Rank and Range ops if ndims is known.
46
+ if x_rank:
47
+ return tf.constant(np.arange(x_rank, dtype=np.int32))
48
+ else:
49
+ # Otherwise, we rely on Range and Rank to do the right thing at run-time.
50
+ return tf.range(0, tf.rank(x))
51
+
52
+ def standardize_images(
53
+ images: Union[tf.Tensor, np.ndarray],
54
+ axis: Union[int, tuple, None] = None
55
+ ) -> tf.Tensor:
56
+ """Standardize image data to have zero mean and unit variance.
57
+
58
+ Args:
59
+ images: The image data.
60
+ axis: The dimensions to standardize across. Exclude the axes that should be
61
+ treated separately, e.g. the channel axis to treat each channel separately.
62
+ If None, standardize across all dimensions.
63
+ Returns:
64
+ images: The standardized image data as tf.float32.
65
+ """
66
+ assert isinstance(images, (tf.Tensor, np.ndarray))
67
+ assert axis is None or isinstance(axis, int) or isinstance(axis, tuple)
68
+ # Convert to tf.Tensor if necessary
69
+ if not tf.is_tensor(images):
70
+ images = tf.convert_to_tensor(images)
71
+ # We are working with tf.float32
72
+ images = tf.cast(images, dtype=tf.float32)
73
+ # Resolve axis arg
74
+ axis = _reduction_dims(images, axis)
75
+ # Compute the mean and std
76
+ num_pixels = tf.math.reduce_prod(tf.gather(tf.shape(images), axis))
77
+ mean = tf.math.reduce_mean(images, axis, keepdims=True)
78
+ std = tf.math.reduce_std(images, axis, keepdims=True)
79
+ # Apply a minimum normalization that protects us against uniform images
80
+ min_std = tf.math.rsqrt(tf.cast(num_pixels, dtype=tf.float32))
81
+ # Perform standardization
82
+ images = tf.subtract(images, mean)
83
+ images = tf.divide(images, tf.maximum(std, min_std))
84
+ return images
85
+
86
+ def normalize_images(
87
+ images: Union[tf.Tensor, np.ndarray],
88
+ axis: Union[int, tuple, None] = None
89
+ ) -> tf.Tensor:
90
+ """Normalize image data to have zero mean.
91
+
92
+ Args:
93
+ images: The image data.
94
+ axis: The dimensions to normalize across. Exclude the axes that should be
95
+ treated separately, e.g. the channel axis to treat each channel separately.
96
+ If None, normalize across all dimensions.
97
+ Returns:
98
+ images: The normalized image data.
99
+ """
100
+ assert isinstance(images, (tf.Tensor, np.ndarray))
101
+ assert axis is None or isinstance(axis, int) or isinstance(axis, tuple)
102
+ # Convert to tf.Tensor if necessary
103
+ if not tf.is_tensor(images):
104
+ images = tf.convert_to_tensor(images)
105
+ # We are working with tf.float32
106
+ images = tf.cast(images, dtype=tf.float32)
107
+ # Resolve axis arg
108
+ axis = _reduction_dims(images, axis)
109
+ # Compute the mean
110
+ mean = tf.math.reduce_mean(images, axis, keepdims=True)
111
+ # Perform normalization
112
+ images = tf.subtract(images, mean)
113
+ return images
114
+
115
+ def normalized_image_diff(
116
+ images: Union[tf.Tensor, np.ndarray],
117
+ axis: int = 0
118
+ ) -> tf.Tensor:
119
+ """Compute the normalized difference of adjacent images.
120
+
121
+ Args:
122
+ images: The image data as float32 in range [0, 1]
123
+ axis: Scalar, the dimension across which to calculate difference
124
+ normalization (e.g., the temporal/sequence dimension).
125
+ Returns:
126
+ images: The processed image data.
127
+ """
128
+ assert isinstance(images, (tf.Tensor, np.ndarray))
129
+ assert axis==0 or axis==1, "Only axis=0 or axis=1 supported"
130
+ # Convert to tf.Tensor if necessary
131
+ if not tf.is_tensor(images):
132
+ images = tf.convert_to_tensor(images)
133
+ # We are working with tf.float32
134
+ images = tf.cast(images, dtype=tf.float32)
135
+ diff = tf.cond(tf.equal(axis, 0),
136
+ true_fn=lambda: images[1:] - images[:-1],
137
+ false_fn=lambda: images[:,1:] - images[:,:-1])
138
+ sum = tf.cond(tf.equal(axis, 0),
139
+ true_fn=lambda: images[1:] + images[:-1],
140
+ false_fn=lambda: images[:,1:] + images[:,:-1])
141
+ sum = tf.clip_by_value(sum, clip_value_min=1e-7, clip_value_max=2)
142
+ return diff / sum
143
+
144
+ def display_scale(
145
+ image: tf.Tensor
146
+ ) -> tf.Tensor:
147
+ """Scale any float32 image to the [0, 1] range for display.
148
+
149
+ Args:
150
+ image: The image tensor.
151
+ Returns:
152
+ out: The scaled image for display.
153
+ """
154
+ assert isinstance(image, tf.Tensor)
155
+ min = tf.math.reduce_min(image)
156
+ range = tf.math.reduce_max(image) - min
157
+ return tf.clip_by_value((image-min) * 1.0/range, 0, 1)
158
+
159
+ def resize_with_random_method(
160
+ images: tf.Tensor,
161
+ target_shape: Union[tuple, None] = (640, 640)
162
+ ) -> tf.Tensor:
163
+ """Resize image(s) with a random method.
164
+
165
+ Args:
166
+ images: Image data, either shape (b, h, w, c) or (h, w, c)
167
+ target_shape: The resize shape in form (new_h, new_w)
168
+ Returns:
169
+ out: The resized image data with shape (b, new_h, new_w, c) or (new_h, new_w, c)
170
+ """
171
+ assert isinstance(images, tf.Tensor)
172
+ assert target_shape is None or (isinstance(target_shape, tuple) and len(target_shape) == 2 and all(isinstance(i, int) for i in target_shape))
173
+ # Draw a random number to determine resize method
174
+ resize_method = tf.random.uniform([], 0, 5, dtype=tf.int32)
175
+ def resize(method):
176
+ def _resize():
177
+ return tf.image.resize(
178
+ images, target_shape, method=method, antialias=True, preserve_aspect_ratio=False)
179
+ return _resize
180
+ # Resize using a random method
181
+ images = tf.case([(tf.equal(resize_method, 0), resize('bicubic')),
182
+ (tf.equal(resize_method, 1), resize('area')),
183
+ (tf.equal(resize_method, 2), resize('nearest')),
184
+ (tf.equal(resize_method, 3), resize('lanczos3'))],
185
+ default=resize('bilinear'))
186
+ return images
187
+
188
+ def random_distortion(
189
+ images: tf.Tensor
190
+ ) -> tf.Tensor:
191
+ """Apply random distortion to image(s)
192
+
193
+ Args:
194
+ images: The image data (b, h, w, c) or (h, w, c)
195
+ Returns:
196
+ images: The distorted image data (b, h, w, c) or (h, w, c)
197
+ """
198
+ assert isinstance(images, tf.Tensor)
199
+ images = tf.image.random_brightness(images, 0.4)
200
+ images = tf.image.random_contrast(images, 0.5, 1.5)
201
+ images = tf.image.random_saturation(images, 0.5, 1.5)
202
+ images = tf.image.random_hue(images, 0.1)
203
+ return images
@@ -0,0 +1,104 @@
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 tensorflow as tf
22
+
23
+ def balanced_sample_weights(
24
+ labels: tf.Tensor,
25
+ unique: tf.Tensor
26
+ ) -> tf.Tensor:
27
+ """Calculate weights for a batch of dense categorical labels intended to be
28
+ multiplied with the losses
29
+
30
+ - Larger weights for examples of under-represented classes, and smaller weights
31
+ for overrepresented classes, while keeping the total loss constant.
32
+ - Important: Only works for dense label representation that equal range from 0 to n
33
+
34
+ Args:
35
+ labels: The dense categorical labels of shape (batch_size,) or (batch_size, 1)
36
+ unique: The unique labels of shape (n_unique_labels,)
37
+ Returns:
38
+ weights: The weights with same shape as labels.
39
+ """
40
+ assert isinstance(labels, tf.Tensor)
41
+ assert isinstance(unique, tf.Tensor)
42
+ # Remove empty dim if necessary
43
+ f_labels = tf.cast(tf.squeeze(labels), tf.int32)
44
+ # Determine count of labels in batch
45
+ # Cannot use bincount to be compatible with XLA
46
+ # count = tf.math.bincount(f_labels, minlength=n_labels)
47
+ def count(x):
48
+ return tf.reduce_sum(tf.cast(tf.equal(x, f_labels), tf.int32))
49
+ count = tf.map_fn(fn=count, elems=unique, fn_output_signature=tf.int32)
50
+ # Batch size and number of unique labels
51
+ batch_size = tf.size(f_labels)
52
+ unique_count = tf.reduce_sum(tf.cast(tf.math.greater(count, 0), tf.int32))
53
+ # Calculate the weight for each class
54
+ class_weights = tf.math.divide_no_nan(tf.cast(batch_size, tf.float32), tf.cast(count, tf.float32))
55
+ class_weights = class_weights / tf.cast(unique_count, tf.float32)
56
+ # Gather weights according to the actual categories of the batch elements
57
+ sample_weights = tf.gather(class_weights, f_labels)
58
+ # Reshape to original shape
59
+ sample_weights = tf.reshape(sample_weights, tf.shape(labels))
60
+ return sample_weights
61
+
62
+ def smooth_l1_loss(
63
+ y_true: tf.Tensor,
64
+ y_pred: tf.Tensor,
65
+ keepdims: bool = False
66
+ ) -> tf.Tensor:
67
+ """Smooth L1 loss
68
+
69
+ Args:
70
+ y_true: Labels. Shape arbitrary.
71
+ y_pred: Predictions. Shape same as y_true.
72
+ keepdims: Keep original shape? Otherwise return global mean.
73
+ Returns:
74
+ loss: The loss. Shape same as original if keepdims, otherwise ()
75
+ """
76
+ assert isinstance(y_true, tf.Tensor)
77
+ assert isinstance(y_pred, tf.Tensor)
78
+ assert isinstance(keepdims, bool)
79
+ t = tf.abs(y_pred - y_true)
80
+ loss = tf.where(t < 1, 0.5 * t ** 2, t - 0.5)
81
+ return tf.cond(tf.equal(keepdims, tf.constant(True)),
82
+ true_fn=lambda: loss,
83
+ false_fn=lambda: tf.reduce_mean(loss))
84
+
85
+ def mae_loss(
86
+ y_true: tf.Tensor,
87
+ y_pred: tf.Tensor,
88
+ keepdims: bool = False
89
+ ) -> tf.Tensor:
90
+ """Mean absolute error loss
91
+
92
+ Args:
93
+ y_true: Labels. Shape arbitrary.
94
+ y_pred: Predictions. Shape same as y_true.
95
+ keepdims: Keep batch dim? Otherwise return global mean.
96
+ Returns:
97
+ loss: The loss. Shape: If keepdims, original shape without final dim, otherwise ()
98
+ """
99
+ assert isinstance(y_true, tf.Tensor)
100
+ assert isinstance(y_pred, tf.Tensor)
101
+ assert isinstance(keepdims, bool)
102
+ return tf.cond(tf.equal(keepdims, tf.constant(True)),
103
+ true_fn=lambda: tf.reduce_mean(tf.abs(y_pred - y_true), axis=-1),
104
+ false_fn=lambda: tf.reduce_mean(tf.abs(y_pred - y_true)))
@@ -0,0 +1,74 @@
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 tensorflow as tf
22
+ from typing import Union
23
+
24
+ class PiecewiseConstantDecayWithWarmup(tf.keras.optimizers.schedules.LearningRateSchedule):
25
+ """Piecewise constant decay with warmup."""
26
+ def __init__(
27
+ self,
28
+ boundaries: list,
29
+ values: list,
30
+ warmup_init_lr: Union[float, int],
31
+ warmup_steps: Union[int, tf.Variable],
32
+ name: Union[str, None] = None
33
+ ):
34
+ super(PiecewiseConstantDecayWithWarmup, self).__init__()
35
+ assert isinstance(boundaries, list)
36
+ assert isinstance(values, list)
37
+ assert isinstance(warmup_init_lr, (float, int))
38
+ assert isinstance(warmup_steps, (int, tf.Variable))
39
+ if len(boundaries) != len(values) - 1:
40
+ raise ValueError("The length of boundaries should be 1 less than the length of values")
41
+ self.boundaries = boundaries
42
+ self.values = values
43
+ self.name = name
44
+ self.warmup_steps = warmup_steps
45
+ self.warmup_init_lr = warmup_init_lr
46
+ def __call__(
47
+ self,
48
+ step: int
49
+ ) -> float:
50
+ assert isinstance(step, (int, tf.Variable))
51
+ with tf.name_scope(self.name or "PiecewiseConstantWarmUp"):
52
+ step = tf.cast(tf.convert_to_tensor(step), tf.float32)
53
+ pred_fn_pairs = []
54
+ warmup_steps = self.warmup_steps
55
+ boundaries = self.boundaries
56
+ values = self.values
57
+ warmup_init_lr = self.warmup_init_lr
58
+ pred_fn_pairs.append((step <= warmup_steps, lambda: warmup_init_lr + step * (values[0] - warmup_init_lr) / warmup_steps))
59
+ pred_fn_pairs.append((tf.logical_and(step <= boundaries[0], step > warmup_steps), lambda: tf.constant(values[0])))
60
+ pred_fn_pairs.append((step > boundaries[-1], lambda: tf.constant(values[-1])))
61
+ for low, high, v in zip(boundaries[:-1], boundaries[1:], values[1:-1]):
62
+ pred = (step > low) & (step <= high)
63
+ pred_fn_pairs.append((pred, lambda v=v: tf.constant(v)))
64
+ # The default isn't needed here because our conditions are mutually
65
+ # exclusive and exhaustive, but tf.case requires it.
66
+ return tf.case(pred_fn_pairs, lambda: tf.constant(values[0]), exclusive=True)
67
+ def get_config(self):
68
+ return {
69
+ "boundaries": self.boundaries,
70
+ "values": self.values,
71
+ "warmup_steps": self.warmup_steps,
72
+ "warmup_init_lr": self.warmup_init_lr,
73
+ "name": self.name
74
+ }