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,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()