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/__init__.py +19 -0
- prpy/constants.py +24 -0
- prpy/ffmpeg/__init__.py +19 -0
- prpy/ffmpeg/probe.py +86 -0
- prpy/ffmpeg/readwrite.py +412 -0
- prpy/ffmpeg/utils.py +81 -0
- prpy/numpy/__init__.py +19 -0
- prpy/numpy/face.py +213 -0
- prpy/numpy/image.py +141 -0
- prpy/numpy/metric.py +179 -0
- prpy/numpy/signal.py +649 -0
- prpy/numpy/stride_tricks.py +182 -0
- prpy/tensorflow/__init__.py +19 -0
- prpy/tensorflow/image.py +203 -0
- prpy/tensorflow/loss.py +104 -0
- prpy/tensorflow/lr_schedule.py +74 -0
- prpy/tensorflow/model_saver.py +203 -0
- prpy/tensorflow/nan.py +226 -0
- prpy/tensorflow/optimizer.py +233 -0
- prpy/tensorflow/signal.py +103 -0
- prpy/torch/__init__.py +19 -0
- prpy/torch/model_saver.py +208 -0
- prpy-0.2.2.dist-info/LICENSE +19 -0
- prpy-0.2.2.dist-info/METADATA +75 -0
- prpy-0.2.2.dist-info/RECORD +27 -0
- prpy-0.2.2.dist-info/WHEEL +5 -0
- prpy-0.2.2.dist-info/top_level.txt +1 -0
|
@@ -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
|
+
"""Save the best models"""
|
|
22
|
+
|
|
23
|
+
import glob
|
|
24
|
+
import logging
|
|
25
|
+
import os
|
|
26
|
+
import shutil
|
|
27
|
+
import tensorflow as tf
|
|
28
|
+
from typing import Union, Callable
|
|
29
|
+
|
|
30
|
+
class Candidate(object):
|
|
31
|
+
"""A candidate model with a score"""
|
|
32
|
+
def __init__(
|
|
33
|
+
self,
|
|
34
|
+
score: Union[float, int],
|
|
35
|
+
dir: str,
|
|
36
|
+
filename: str
|
|
37
|
+
):
|
|
38
|
+
"""Initialise a new candidate.
|
|
39
|
+
|
|
40
|
+
Args:
|
|
41
|
+
score: The score achieved by the candidate
|
|
42
|
+
dir: The directory where the candidate model can be saved
|
|
43
|
+
filename: The filename under which the candidate model can be saved
|
|
44
|
+
"""
|
|
45
|
+
assert isinstance(score, (float, int))
|
|
46
|
+
assert isinstance(dir, str)
|
|
47
|
+
assert isinstance(filename, str)
|
|
48
|
+
self.score = score
|
|
49
|
+
self.filepath = os.path.join(dir, filename)
|
|
50
|
+
|
|
51
|
+
class ModelSaver(object):
|
|
52
|
+
"""Save the best models to disk"""
|
|
53
|
+
def __init__(
|
|
54
|
+
self,
|
|
55
|
+
dir: str = "checkpoints",
|
|
56
|
+
keep_best: int = 5,
|
|
57
|
+
keep_latest: int = 1,
|
|
58
|
+
save_format: str = "tf",
|
|
59
|
+
save_optimizer: bool = False,
|
|
60
|
+
compare_fn: Callable[[float, float], bool] = lambda x,y: x.score < y.score,
|
|
61
|
+
sort_reverse: bool = False,
|
|
62
|
+
log_fn: Callable[[str], None] = logging.info
|
|
63
|
+
):
|
|
64
|
+
"""Init the ModelSaver
|
|
65
|
+
|
|
66
|
+
Args:
|
|
67
|
+
dir: The directory where models should be saved
|
|
68
|
+
keep_best: The number of best scoring models to keep
|
|
69
|
+
keep_latest: The number of latest models to keep
|
|
70
|
+
save_format: Model format for saving ['tf' or 'h5']
|
|
71
|
+
save_optimizer: Also save optimizer state?
|
|
72
|
+
compare_fn: Function that compares two scores
|
|
73
|
+
sort_reverse: Reverse sort order?
|
|
74
|
+
log_fn: Function to write logs
|
|
75
|
+
"""
|
|
76
|
+
assert isinstance(dir, str)
|
|
77
|
+
assert isinstance(keep_best, int)
|
|
78
|
+
assert isinstance(keep_latest, int)
|
|
79
|
+
assert isinstance(save_format, str)
|
|
80
|
+
assert isinstance(save_optimizer, bool)
|
|
81
|
+
assert callable(compare_fn)
|
|
82
|
+
assert isinstance(sort_reverse, bool)
|
|
83
|
+
assert callable(log_fn)
|
|
84
|
+
self.best_candidates = []
|
|
85
|
+
self.latest_candidates = []
|
|
86
|
+
# The destination directory (make if necessary)
|
|
87
|
+
self.dir = dir
|
|
88
|
+
if not os.path.exists(self.dir):
|
|
89
|
+
os.makedirs(self.dir)
|
|
90
|
+
self.keep_best = keep_best
|
|
91
|
+
self.keep_latest = keep_latest
|
|
92
|
+
self.save_format = save_format
|
|
93
|
+
self.save_optimizer = save_optimizer
|
|
94
|
+
self.compare_fn_best = compare_fn
|
|
95
|
+
self.compare_fn_latest = lambda x,y: x.score > y.score
|
|
96
|
+
self.sort_reverse = sort_reverse
|
|
97
|
+
self.log_fn = log_fn
|
|
98
|
+
|
|
99
|
+
def __save(self, model: tf.keras.Model, filepath: str):
|
|
100
|
+
"""Save a model to disk.
|
|
101
|
+
|
|
102
|
+
Args:
|
|
103
|
+
model: The keras model to be saved
|
|
104
|
+
filepath: The filepath to save to
|
|
105
|
+
"""
|
|
106
|
+
assert isinstance(model, tf.keras.Model)
|
|
107
|
+
assert isinstance(filepath, str)
|
|
108
|
+
if self.save_format == 'h5': filepath += '.h5'
|
|
109
|
+
# Save model
|
|
110
|
+
if self.save_optimizer:
|
|
111
|
+
model.save(
|
|
112
|
+
filepath=filepath, overwrite=True, include_optimizer=True,
|
|
113
|
+
save_format=self.save_format)
|
|
114
|
+
else:
|
|
115
|
+
model.save_weights(
|
|
116
|
+
filepath=filepath, overwrite=True, save_format=self.save_format)
|
|
117
|
+
|
|
118
|
+
def save_keep(self, model: tf.keras.Model, step: int, name: str):
|
|
119
|
+
"""Save and keep the given model.
|
|
120
|
+
|
|
121
|
+
Args:
|
|
122
|
+
model: The model to be saved and kept
|
|
123
|
+
step: The current training step
|
|
124
|
+
name: The model name
|
|
125
|
+
"""
|
|
126
|
+
assert isinstance(model, tf.keras.Model)
|
|
127
|
+
assert isinstance(step, int)
|
|
128
|
+
assert isinstance(name, str)
|
|
129
|
+
self.log_fn("Saving and keeping model for step {}".format(step))
|
|
130
|
+
filepath = os.path.join(self.dir, name + "_keep_" + str(step))
|
|
131
|
+
self.__save(model=model, filepath=filepath)
|
|
132
|
+
|
|
133
|
+
def save_latest(self, model: tf.keras.Model, step: int, name: str):
|
|
134
|
+
"""Save the given model as currently latest.
|
|
135
|
+
|
|
136
|
+
Args:
|
|
137
|
+
model: The model to be saved as latest
|
|
138
|
+
step: The current training step
|
|
139
|
+
name: The model name
|
|
140
|
+
"""
|
|
141
|
+
assert isinstance(model, tf.keras.Model)
|
|
142
|
+
assert isinstance(step, int)
|
|
143
|
+
assert isinstance(name, str)
|
|
144
|
+
name = name + "_latest_" + str(step)
|
|
145
|
+
# Use step as score
|
|
146
|
+
candidate = Candidate(score=step, dir=self.dir, filename=name)
|
|
147
|
+
if len(self.latest_candidates) < self.keep_latest \
|
|
148
|
+
or self.compare_fn_latest(candidate, self.latest_candidates[-1]):
|
|
149
|
+
self.log_fn("Saving latest model for step {}".format(step))
|
|
150
|
+
# Keep candidate
|
|
151
|
+
self.latest_candidates.append(candidate)
|
|
152
|
+
self.latest_candidates = sorted(
|
|
153
|
+
self.latest_candidates, key=lambda x: x.score, reverse=True)
|
|
154
|
+
# Save candidate
|
|
155
|
+
self.__save(model, filepath=candidate.filepath)
|
|
156
|
+
# Prune candidate
|
|
157
|
+
for candidate in self.latest_candidates[self.keep_latest:]:
|
|
158
|
+
for file in glob.glob(r'{}*'.format(candidate.filepath)):
|
|
159
|
+
if self.save_format == 'tf' and self.save_optimizer:
|
|
160
|
+
shutil.rmtree(file)
|
|
161
|
+
else:
|
|
162
|
+
os.remove(file)
|
|
163
|
+
self.latest_candidates = self.latest_candidates[0:self.keep_latest]
|
|
164
|
+
|
|
165
|
+
def save_best(self, model: tf.keras.Model, score: float, step: int, name: str):
|
|
166
|
+
"""Save the given model as a candidate for best model.
|
|
167
|
+
|
|
168
|
+
Args:
|
|
169
|
+
model: The model to be saved as candidate for best model
|
|
170
|
+
score: The score achieved by the model
|
|
171
|
+
step: The current training step
|
|
172
|
+
name: The name of the model
|
|
173
|
+
"""
|
|
174
|
+
assert isinstance(model, tf.keras.Model)
|
|
175
|
+
assert isinstance(score, float)
|
|
176
|
+
assert isinstance(step, int)
|
|
177
|
+
assert isinstance(name, str)
|
|
178
|
+
self.log_fn('Saving model for step {}'.format(step))
|
|
179
|
+
name = name + "_best_" + str(step)
|
|
180
|
+
candidate = Candidate(score, dir=self.dir, filename=name)
|
|
181
|
+
if len(self.best_candidates) < self.keep_best \
|
|
182
|
+
or self.compare_fn_best(candidate, self.best_candidates[-1]):
|
|
183
|
+
# Keep candidate
|
|
184
|
+
self.log_fn("Keeping model {} with score {:.4f}".format(
|
|
185
|
+
candidate.filepath, candidate.score))
|
|
186
|
+
self.best_candidates.append(candidate)
|
|
187
|
+
self.best_candidates = sorted(
|
|
188
|
+
self.best_candidates, key=lambda x: x.score, reverse=self.sort_reverse)
|
|
189
|
+
# Save candidate
|
|
190
|
+
self.__save(model, filepath=candidate.filepath)
|
|
191
|
+
# Prune candidates
|
|
192
|
+
for candidate in self.best_candidates[self.keep_best:]:
|
|
193
|
+
self.log_fn('Removing old model {} with score {:.4f}'.format(
|
|
194
|
+
candidate.filepath, candidate.score))
|
|
195
|
+
for file in glob.glob(r'{}*'.format(candidate.filepath)):
|
|
196
|
+
if self.save_format == 'tf' and self.save_optimizer:
|
|
197
|
+
shutil.rmtree(file)
|
|
198
|
+
else:
|
|
199
|
+
os.remove(file)
|
|
200
|
+
self.best_candidates = self.best_candidates[0:self.keep_best]
|
|
201
|
+
else:
|
|
202
|
+
# Skip the candidate
|
|
203
|
+
self.log_fn('Skipping candidate {}'.format(candidate.filepath))
|
prpy/tensorflow/nan.py
ADDED
|
@@ -0,0 +1,226 @@
|
|
|
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, Tuple, Callable
|
|
23
|
+
|
|
24
|
+
def reduce_nanmean(
|
|
25
|
+
x: tf.Tensor,
|
|
26
|
+
axis: Union[int, tuple, None] = None
|
|
27
|
+
) -> tf.Tensor:
|
|
28
|
+
"""tf.reduce_mean, ignoring non-finite vals.
|
|
29
|
+
|
|
30
|
+
- Returns `nan` for all-nan slices.
|
|
31
|
+
|
|
32
|
+
Args:
|
|
33
|
+
x: The input tensor.
|
|
34
|
+
axis: The dimension to reduce.
|
|
35
|
+
Returns:
|
|
36
|
+
The reduced tensor.
|
|
37
|
+
"""
|
|
38
|
+
assert isinstance(x, tf.Tensor)
|
|
39
|
+
assert axis is None or isinstance(axis, (int, tuple))
|
|
40
|
+
mask = tf.math.is_finite(x)
|
|
41
|
+
numerator = tf.reduce_sum(tf.where(mask, x, tf.zeros_like(x)), axis=axis)
|
|
42
|
+
denominator = tf.reduce_sum(tf.cast(mask, dtype=x.dtype), axis=axis)
|
|
43
|
+
return numerator / denominator
|
|
44
|
+
|
|
45
|
+
class ReduceNanMean:
|
|
46
|
+
"""tf.reduce_mean, ignoring non-finite values. Supports gradient.
|
|
47
|
+
|
|
48
|
+
Behavior when x is non-finite:
|
|
49
|
+
- out: Non-finite vals in a slice contribute 0
|
|
50
|
+
- out: All-non-finite slices are nan
|
|
51
|
+
- grad = 0
|
|
52
|
+
"""
|
|
53
|
+
def __init__(self, axis: Union[int, tuple, None] = None):
|
|
54
|
+
"""Initialize.
|
|
55
|
+
|
|
56
|
+
Args:
|
|
57
|
+
axis: Axes to reduce by mean
|
|
58
|
+
"""
|
|
59
|
+
assert axis is None or isinstance(axis, int) or (isinstance(axis, tuple) and all(isinstance(i, int) for i in axis))
|
|
60
|
+
self.axis = axis
|
|
61
|
+
@tf.custom_gradient
|
|
62
|
+
def __call__(self, x: tf.Tensor) -> Tuple[tf.Tensor, Callable[[tf.Tensor], tf.Tensor]]:
|
|
63
|
+
"""Compute the mean.
|
|
64
|
+
|
|
65
|
+
Args:
|
|
66
|
+
x: The values.
|
|
67
|
+
Returns:
|
|
68
|
+
out: The computed mean.
|
|
69
|
+
grad: Function calculating the gradient
|
|
70
|
+
"""
|
|
71
|
+
assert isinstance(x, tf.Tensor)
|
|
72
|
+
mask = tf.math.is_finite(x)
|
|
73
|
+
num = tf.reduce_sum(tf.where(mask, x, tf.zeros_like(x)), axis=self.axis)
|
|
74
|
+
den = tf.reduce_sum(tf.cast(mask, dtype=x.dtype), axis=self.axis)
|
|
75
|
+
mean = num / den
|
|
76
|
+
def grad(upstream: tf.Tensor) -> tf.Tensor:
|
|
77
|
+
den = tf.reduce_sum(tf.cast(mask, dtype=x.dtype), axis=self.axis)
|
|
78
|
+
# Tile upstream to match x
|
|
79
|
+
if self.axis is not None:
|
|
80
|
+
axis_list = list(self.axis) if isinstance(self.axis, tuple) else [self.axis]
|
|
81
|
+
# Expand dims to match x
|
|
82
|
+
for axis in axis_list:
|
|
83
|
+
upstream = tf.expand_dims(upstream, axis=axis)
|
|
84
|
+
den = tf.expand_dims(den, axis=axis)
|
|
85
|
+
# Tile
|
|
86
|
+
tile_shape = [1 if s1 == s2 else s2 if s1 == 1 else s1 for s1, s2 in zip(x.shape, upstream.shape)]
|
|
87
|
+
upstream = tf.tile(upstream, tile_shape)
|
|
88
|
+
den = tf.tile(den, tile_shape)
|
|
89
|
+
# Compute gradient and set to 0 where input was not finite
|
|
90
|
+
dout_dx = tf.where(mask, upstream / den, tf.zeros_like(x))
|
|
91
|
+
return dout_dx
|
|
92
|
+
return mean, grad
|
|
93
|
+
|
|
94
|
+
def reduce_nansum(
|
|
95
|
+
x: tf.Tensor,
|
|
96
|
+
weight: Union[tf.Tensor, list, None] = None,
|
|
97
|
+
axis: Union[int, tuple, None] = None,
|
|
98
|
+
default: Union[float, int] = float('nan')
|
|
99
|
+
) -> tf.Tensor:
|
|
100
|
+
"""tf.reduce_sum, weighted by weight, ignoring non-finite values.
|
|
101
|
+
|
|
102
|
+
- Returns default for all-nan slices.
|
|
103
|
+
|
|
104
|
+
Args:
|
|
105
|
+
x: The input tensor.
|
|
106
|
+
weight: The weight tensor, with the same shape as x or broadcastable to it.
|
|
107
|
+
axis: The dimension to reduce.
|
|
108
|
+
default: The value to return for all-non-finite slices.
|
|
109
|
+
Returns:
|
|
110
|
+
The reduced tensor.
|
|
111
|
+
"""
|
|
112
|
+
assert isinstance(x, tf.Tensor)
|
|
113
|
+
assert weight is None or isinstance(weight, (tf.Tensor, list))
|
|
114
|
+
assert axis is None or isinstance(axis, (int, tuple))
|
|
115
|
+
assert isinstance(default, (float, int))
|
|
116
|
+
mask = tf.math.is_finite(x)
|
|
117
|
+
if weight is None:
|
|
118
|
+
sum = tf.reduce_sum(tf.where(mask, x, tf.zeros_like(x)), axis=axis)
|
|
119
|
+
else:
|
|
120
|
+
weight = tf.where(mask, weight, tf.zeros_like(weight))
|
|
121
|
+
sum = tf.reduce_sum(tf.where(mask, x * tf.cast(weight, x.dtype), tf.zeros_like(x)), axis=axis)
|
|
122
|
+
# If there are no finite elements in a slice, return the default.
|
|
123
|
+
return tf.where(tf.reduce_all(tf.logical_not(mask), axis=axis),
|
|
124
|
+
tf.constant(default, dtype=x.dtype),
|
|
125
|
+
sum)
|
|
126
|
+
|
|
127
|
+
class ReduceNanSum:
|
|
128
|
+
"""tf.reduce_sum, weighted by weight, ignoring non-finite values. Supports gradient.
|
|
129
|
+
|
|
130
|
+
Behavior when x is non-finite:
|
|
131
|
+
- out: Non-finite vals in a slice contribute 0
|
|
132
|
+
- out: All-non-finite slices are set to default value
|
|
133
|
+
- grad = 0
|
|
134
|
+
"""
|
|
135
|
+
def __init__(
|
|
136
|
+
self,
|
|
137
|
+
weight: Union[tf.Tensor, list, None] = None,
|
|
138
|
+
axis: Union[int, tuple, None] = None,
|
|
139
|
+
default: Union[int, float] = float('nan')
|
|
140
|
+
):
|
|
141
|
+
"""Initialize.
|
|
142
|
+
|
|
143
|
+
Args:
|
|
144
|
+
weight: The weight tensor, with the same shape as x or broadcastable to it.
|
|
145
|
+
axis: Axes to reduce by sum
|
|
146
|
+
default: Value for all-non-finite slices
|
|
147
|
+
"""
|
|
148
|
+
assert weight is None or isinstance(weight, (tf.Tensor, list))
|
|
149
|
+
assert axis is None or isinstance(axis, (int, tuple))
|
|
150
|
+
assert isinstance(default, (int, float))
|
|
151
|
+
self.weight = weight
|
|
152
|
+
self.axis = axis
|
|
153
|
+
self.default = default
|
|
154
|
+
@tf.custom_gradient
|
|
155
|
+
def __call__(self, x: tf.Tensor) -> Tuple[tf.Tensor, Callable[[tf.Tensor], tf.Tensor]]:
|
|
156
|
+
"""Compute the sum.
|
|
157
|
+
|
|
158
|
+
Args:
|
|
159
|
+
x: The values.
|
|
160
|
+
Returns:
|
|
161
|
+
out: The computed sum.
|
|
162
|
+
grad: Function calculating the gradient
|
|
163
|
+
"""
|
|
164
|
+
assert isinstance(x, tf.Tensor)
|
|
165
|
+
mask = tf.math.is_finite(x)
|
|
166
|
+
if self.weight is None:
|
|
167
|
+
sum = tf.reduce_sum(tf.where(mask, x, tf.zeros_like(x)), axis=self.axis)
|
|
168
|
+
else:
|
|
169
|
+
weight = tf.where(mask, self.weight, tf.zeros_like(self.weight))
|
|
170
|
+
sum = tf.reduce_sum(tf.where(mask, x * tf.cast(weight, x.dtype), tf.zeros_like(x)), axis=self.axis)
|
|
171
|
+
# If there are no finite elements in a slice, return the default.
|
|
172
|
+
out = tf.where(tf.reduce_all(tf.logical_not(mask), axis=self.axis),
|
|
173
|
+
tf.cast(self.default, x.dtype),
|
|
174
|
+
sum)
|
|
175
|
+
def grad(upstream: tf.Tensor) -> tf.Tensor:
|
|
176
|
+
# Tile upstream to match x
|
|
177
|
+
if self.axis is not None:
|
|
178
|
+
axis_list = list(self.axis) if isinstance(self.axis, tuple) else [self.axis]
|
|
179
|
+
# Expand dims to match x
|
|
180
|
+
for axis in axis_list:
|
|
181
|
+
upstream = tf.expand_dims(upstream, axis=axis)
|
|
182
|
+
# Tile
|
|
183
|
+
tile_shape = [1 if s1 == s2 else s2 if s1 == 1 else s1 for s1, s2 in zip(x.shape, upstream.shape)]
|
|
184
|
+
upstream = tf.tile(upstream, tile_shape)
|
|
185
|
+
# Set the gradient to 0 where input was not finite
|
|
186
|
+
dout_dx = tf.where(mask, upstream, tf.zeros_like(x))
|
|
187
|
+
return dout_dx
|
|
188
|
+
return out, grad
|
|
189
|
+
|
|
190
|
+
class NanLinearCombination:
|
|
191
|
+
"""Linear combination with one fixed element. Supports gradient.
|
|
192
|
+
|
|
193
|
+
- Calculates (1 - x) * val_1 + x * val_2
|
|
194
|
+
- Behavior when val_1 or val_2 are non-finite:
|
|
195
|
+
- out = nan
|
|
196
|
+
- grad = 0
|
|
197
|
+
"""
|
|
198
|
+
@tf.custom_gradient
|
|
199
|
+
def __call__(
|
|
200
|
+
self,
|
|
201
|
+
x: tf.Tensor,
|
|
202
|
+
val_1: tf.Tensor,
|
|
203
|
+
val_2: tf.Tensor
|
|
204
|
+
) -> Tuple[tf.Tensor, Callable[[tf.Tensor], Tuple[tf.Tensor, tf.Tensor, tf.Tensor]]]:
|
|
205
|
+
"""Compute the linear combination.
|
|
206
|
+
|
|
207
|
+
Args:
|
|
208
|
+
x: The combination weight. All elements must be in [0, 1]
|
|
209
|
+
val_1: The second value used in the linear combination.
|
|
210
|
+
val_2: The second value used in the linear combination.
|
|
211
|
+
Returns:
|
|
212
|
+
out: The computed sum.
|
|
213
|
+
grad: Function calculating the gradient
|
|
214
|
+
"""
|
|
215
|
+
# Compute the linear combination
|
|
216
|
+
val_1 = tf.broadcast_to(val_1, x.shape)
|
|
217
|
+
val_2 = tf.broadcast_to(val_2, x.shape)
|
|
218
|
+
out = (1 - x) * val_1 + x * val_2
|
|
219
|
+
def grad(upstream: tf.Tensor) -> Tuple[tf.Tensor, tf.Tensor, tf.Tensor]:
|
|
220
|
+
mask = tf.math.logical_and(tf.math.is_finite(val_1), tf.math.is_finite(val_2))
|
|
221
|
+
dout_dx = tf.where(mask, upstream * (val_2 - val_1), tf.zeros_like(x))
|
|
222
|
+
dout_dval_1 = upstream * (1 - x)
|
|
223
|
+
dout_dval_2 = upstream * x
|
|
224
|
+
return dout_dx, dout_dval_1, dout_dval_2
|
|
225
|
+
return out, grad
|
|
226
|
+
|
|
@@ -0,0 +1,233 @@
|
|
|
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
|
+
from packaging import version
|
|
22
|
+
import tensorflow as tf
|
|
23
|
+
|
|
24
|
+
if version.parse(tf.__version__) <= version.parse("2.6.5"):
|
|
25
|
+
from tensorflow.keras.optimizers import Adam # Legacy / V2 optimizer
|
|
26
|
+
from tensorflow.keras.mixed_precision import LossScaleOptimizer
|
|
27
|
+
elif version.parse(tf.__version__) >= version.parse("2.9") and version.parse(tf.__version__) <= version.parse("2.11"):
|
|
28
|
+
from keras.optimizers.optimizer_experimental.adam import Adam # Experimental / V3 optimizer
|
|
29
|
+
from keras.optimizers.optimizer_experimental.adamw import AdamW # Experimental / V3 optimizer
|
|
30
|
+
from keras.mixed_precision.loss_scale_optimizer import LossScaleOptimizerV3 as LossScaleOptimizer
|
|
31
|
+
elif version.parse(tf.__version__) > version.parse("2.11"):
|
|
32
|
+
from tensorflow.keras.optimizers import Adam, AdamW # V3 optimizer
|
|
33
|
+
from keras.mixed_precision.loss_scale_optimizer import LossScaleOptimizerV3 as LossScaleOptimizer
|
|
34
|
+
else:
|
|
35
|
+
raise ImportError("This version of TensorFlow is not compatible.")
|
|
36
|
+
|
|
37
|
+
class EpochAdamMetaclass(type):
|
|
38
|
+
"""Metaclass that delegates EpochAdam instance creation."""
|
|
39
|
+
def __call__(cls, **kwargs):
|
|
40
|
+
if version.parse(tf.__version__) <= version.parse("2.6.5"):
|
|
41
|
+
return EpochAdamV2(**kwargs)
|
|
42
|
+
else:
|
|
43
|
+
return EpochAdamExperimental(**kwargs)
|
|
44
|
+
|
|
45
|
+
class EpochAdamWMetaclass(type):
|
|
46
|
+
"""Metaclass that delegates EpochAdamW instance creation."""
|
|
47
|
+
def __call__(cls, **kwargs):
|
|
48
|
+
if version.parse(tf.__version__) > version.parse("2.9"):
|
|
49
|
+
return EpochAdamWExperimental(**kwargs)
|
|
50
|
+
else:
|
|
51
|
+
raise ImportError("This version of TensorFlow is not compatible.")
|
|
52
|
+
|
|
53
|
+
class EpochAdam(metaclass=EpochAdamMetaclass):
|
|
54
|
+
pass
|
|
55
|
+
|
|
56
|
+
class EpochAdamW(metaclass=EpochAdamWMetaclass):
|
|
57
|
+
pass
|
|
58
|
+
|
|
59
|
+
class EpochAdamExperimental(Adam):
|
|
60
|
+
"""Experimental Adam optimizer that retrieves learning rate based on epochs"""
|
|
61
|
+
def __init__(self, **kwargs):
|
|
62
|
+
# Create epochs counter variable
|
|
63
|
+
with tf.init_scope():
|
|
64
|
+
# Lift the variable creation to init scope to avoid environment issue.
|
|
65
|
+
self._epochs = tf.Variable(
|
|
66
|
+
0, name="epochs", dtype=tf.int64, trainable=False)
|
|
67
|
+
super().__init__(**kwargs)
|
|
68
|
+
self._variables.append(self._epochs)
|
|
69
|
+
@property
|
|
70
|
+
def epochs(self):
|
|
71
|
+
return self._epochs
|
|
72
|
+
@epochs.setter
|
|
73
|
+
def epochs(self, variable):
|
|
74
|
+
if getattr(self, "_built", False):
|
|
75
|
+
raise RuntimeError(
|
|
76
|
+
"Cannot set `epochs` to a new Variable after "
|
|
77
|
+
"the Optimizer weights have been created. Here it is "
|
|
78
|
+
f"attempting to set `iterations` to {variable}."
|
|
79
|
+
"Usually this means you are trying to set `iterations`"
|
|
80
|
+
" after calling `apply_gradients()`. Please set "
|
|
81
|
+
"`iterations` before calling `apply_gradients()`.")
|
|
82
|
+
self._epochs = variable
|
|
83
|
+
def _build_learning_rate(self, learning_rate):
|
|
84
|
+
with tf.init_scope():
|
|
85
|
+
if isinstance(learning_rate, tf.keras.optimizers.schedules.LearningRateSchedule):
|
|
86
|
+
# Create a variable to hold the current learning rate.
|
|
87
|
+
current_learning_rate = tf.convert_to_tensor(
|
|
88
|
+
learning_rate(self.epochs))
|
|
89
|
+
self._current_learning_rate = tf.Variable(
|
|
90
|
+
current_learning_rate,
|
|
91
|
+
name="current_learning_rate",
|
|
92
|
+
dtype=current_learning_rate.dtype,
|
|
93
|
+
trainable=False)
|
|
94
|
+
return learning_rate
|
|
95
|
+
return tf.Variable(
|
|
96
|
+
learning_rate,
|
|
97
|
+
name="learning_rate",
|
|
98
|
+
dtype=tf.float32,
|
|
99
|
+
trainable=False)
|
|
100
|
+
def _compute_current_learning_rate(self):
|
|
101
|
+
if isinstance(self._learning_rate, tf.keras.optimizers.schedules.LearningRateSchedule):
|
|
102
|
+
# Compute the current learning rate at the beginning of variable update.
|
|
103
|
+
if hasattr(self, "_current_learning_rate"):
|
|
104
|
+
self._current_learning_rate.assign(
|
|
105
|
+
self._learning_rate(self.epochs))
|
|
106
|
+
else:
|
|
107
|
+
current_learning_rate = tf.convert_to_tensor(
|
|
108
|
+
self._learning_rate(self.epochs))
|
|
109
|
+
self._current_learning_rate = tf.Variable(
|
|
110
|
+
current_learning_rate,
|
|
111
|
+
name="current_learning_rate",
|
|
112
|
+
dtype=current_learning_rate.dtype,
|
|
113
|
+
trainable=False)
|
|
114
|
+
def finish_epoch(self):
|
|
115
|
+
"""Increment epoch count and re-compute lr"""
|
|
116
|
+
self._epochs.assign_add(1)
|
|
117
|
+
self._compute_current_learning_rate()
|
|
118
|
+
|
|
119
|
+
class EpochAdamWExperimental(AdamW):
|
|
120
|
+
"""Experimental AdamW optimizer that retrieves learning rate based on epochs"""
|
|
121
|
+
def __init__(self, **kwargs):
|
|
122
|
+
# Create epochs counter variable
|
|
123
|
+
with tf.init_scope():
|
|
124
|
+
# Lift the variable creation to init scope to avoid environment issue.
|
|
125
|
+
self._epochs = tf.Variable(
|
|
126
|
+
0, name="epochs", dtype=tf.int64, trainable=False)
|
|
127
|
+
super().__init__(**kwargs)
|
|
128
|
+
self._variables.append(self._epochs)
|
|
129
|
+
@property
|
|
130
|
+
def epochs(self):
|
|
131
|
+
return self._epochs
|
|
132
|
+
@epochs.setter
|
|
133
|
+
def epochs(self, variable):
|
|
134
|
+
if getattr(self, "_built", False):
|
|
135
|
+
raise RuntimeError(
|
|
136
|
+
"Cannot set `epochs` to a new Variable after "
|
|
137
|
+
"the Optimizer weights have been created. Here it is "
|
|
138
|
+
f"attempting to set `iterations` to {variable}."
|
|
139
|
+
"Usually this means you are trying to set `iterations`"
|
|
140
|
+
" after calling `apply_gradients()`. Please set "
|
|
141
|
+
"`iterations` before calling `apply_gradients()`.")
|
|
142
|
+
self._epochs = variable
|
|
143
|
+
def _build_learning_rate(self, learning_rate):
|
|
144
|
+
with tf.init_scope():
|
|
145
|
+
if isinstance(learning_rate, tf.keras.optimizers.schedules.LearningRateSchedule):
|
|
146
|
+
# Create a variable to hold the current learning rate.
|
|
147
|
+
current_learning_rate = tf.convert_to_tensor(
|
|
148
|
+
learning_rate(self.epochs))
|
|
149
|
+
self._current_learning_rate = tf.Variable(
|
|
150
|
+
current_learning_rate,
|
|
151
|
+
name="current_learning_rate",
|
|
152
|
+
dtype=current_learning_rate.dtype,
|
|
153
|
+
trainable=False)
|
|
154
|
+
return learning_rate
|
|
155
|
+
return tf.Variable(
|
|
156
|
+
learning_rate,
|
|
157
|
+
name="learning_rate",
|
|
158
|
+
dtype=tf.float32,
|
|
159
|
+
trainable=False)
|
|
160
|
+
def _compute_current_learning_rate(self):
|
|
161
|
+
if isinstance(self._learning_rate, tf.keras.optimizers.schedules.LearningRateSchedule):
|
|
162
|
+
# Compute the current learning rate at the beginning of variable update.
|
|
163
|
+
if hasattr(self, "_current_learning_rate"):
|
|
164
|
+
self._current_learning_rate.assign(
|
|
165
|
+
self._learning_rate(self.epochs))
|
|
166
|
+
else:
|
|
167
|
+
current_learning_rate = tf.convert_to_tensor(
|
|
168
|
+
self._learning_rate(self.epochs))
|
|
169
|
+
self._current_learning_rate = tf.Variable(
|
|
170
|
+
current_learning_rate,
|
|
171
|
+
name="current_learning_rate",
|
|
172
|
+
dtype=current_learning_rate.dtype,
|
|
173
|
+
trainable=False)
|
|
174
|
+
def finish_epoch(self):
|
|
175
|
+
"""Increment epoch count and re-compute lr"""
|
|
176
|
+
self._epochs.assign_add(1)
|
|
177
|
+
self._compute_current_learning_rate()
|
|
178
|
+
|
|
179
|
+
class EpochAdamV2(Adam):
|
|
180
|
+
"""V2 Adam optimizer that retrieves learning rate based on epochs"""
|
|
181
|
+
def __init__(self, **kwargs):
|
|
182
|
+
super().__init__(**kwargs)
|
|
183
|
+
self._epochs = None
|
|
184
|
+
def _decayed_lr(self, var_dtype):
|
|
185
|
+
"""Get learning rate based on epochs."""
|
|
186
|
+
lr_t = self._get_hyper("learning_rate", var_dtype)
|
|
187
|
+
if isinstance(lr_t, tf.keras.optimizers.schedules.LearningRateSchedule):
|
|
188
|
+
epochs = tf.cast(self.epochs, var_dtype)
|
|
189
|
+
lr_t = tf.cast(lr_t(epochs), var_dtype)
|
|
190
|
+
return lr_t
|
|
191
|
+
@property
|
|
192
|
+
def epochs(self):
|
|
193
|
+
"""Variable. The number of epochs."""
|
|
194
|
+
if self._epochs is None:
|
|
195
|
+
self._epochs = self.add_weight(
|
|
196
|
+
"epochs", shape=[], dtype=tf.int64, trainable=False,
|
|
197
|
+
aggregation=tf.VariableAggregation.ONLY_FIRST_REPLICA)
|
|
198
|
+
self._weights.append(self._epochs)
|
|
199
|
+
return self._epochs
|
|
200
|
+
def finish_epoch(self):
|
|
201
|
+
"""Increment epoch count"""
|
|
202
|
+
return self._epochs.assign_add(1)
|
|
203
|
+
|
|
204
|
+
class EpochLossScaleOptimizerMetaclass(type):
|
|
205
|
+
"""Metaclass that delegates EpochLossScaleOptimizer instance creation."""
|
|
206
|
+
def __call__(cls, inner_optimizer, **kwargs):
|
|
207
|
+
if version.parse(tf.__version__) <= version.parse("2.6.5"):
|
|
208
|
+
return EpochLossScaleOptimizerV2(inner_optimizer, **kwargs)
|
|
209
|
+
else:
|
|
210
|
+
return EpochLossScaleOptimizerV3(inner_optimizer, **kwargs)
|
|
211
|
+
|
|
212
|
+
class EpochLossScaleOptimizer(metaclass=EpochLossScaleOptimizerMetaclass):
|
|
213
|
+
pass
|
|
214
|
+
|
|
215
|
+
class EpochLossScaleOptimizerV2(LossScaleOptimizer):
|
|
216
|
+
"""Subclass LossScaleOptimizer to use epochs. Works for tensorflow<=2.6.5"""
|
|
217
|
+
@property
|
|
218
|
+
def epochs(self):
|
|
219
|
+
"""Variable. The number of epochs."""
|
|
220
|
+
return self._optimizer.epochs
|
|
221
|
+
def finish_epoch(self):
|
|
222
|
+
"""Increment epoch count"""
|
|
223
|
+
return self._optimizer.finish_epoch()
|
|
224
|
+
|
|
225
|
+
class EpochLossScaleOptimizerV3(LossScaleOptimizer):
|
|
226
|
+
"""Subclass LossScaleOptimizer to use epochs."""
|
|
227
|
+
@property
|
|
228
|
+
def epochs(self):
|
|
229
|
+
"""Variable. The number of epochs."""
|
|
230
|
+
return self._optimizer.epochs
|
|
231
|
+
def finish_epoch(self):
|
|
232
|
+
"""Increment epoch count"""
|
|
233
|
+
return self._optimizer.finish_epoch()
|