MCPcopy Create free account
hub / github.com/TPCD/DCCL / IndicatePlateau

Class IndicatePlateau

project_utils/general_utils.py:304–378  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

302
303
304class IndicatePlateau(object):
305
306 def __init__(self, threshold=5e-4, patience_epochs=5, mode='min', threshold_mode='rel'):
307
308 self.patience = patience_epochs
309 self.cooldown_counter = 0
310 self.mode = mode
311 self.threshold = threshold
312 self.threshold_mode = threshold_mode
313 self.best = None
314 self.num_bad_epochs = None
315 self.mode_worse = None # the worse value for the chosen mode
316 self.last_epoch = 0
317 self._init_is_better(mode=mode, threshold=threshold,
318 threshold_mode=threshold_mode)
319
320 self._init_is_better(mode=mode, threshold=threshold,
321 threshold_mode=threshold_mode)
322 self._reset()
323
324 def _reset(self):
325 """Resets num_bad_epochs counter and cooldown counter."""
326 self.best = self.mode_worse
327 self.cooldown_counter = 0
328 self.num_bad_epochs = 0
329
330 def step(self, metrics, epoch=None):
331 # convert `metrics` to float, in case it's a zero-dim Tensor
332 current = float(metrics)
333 self.last_epoch = epoch
334
335 if self.is_better(current, self.best):
336 self.best = current
337 self.num_bad_epochs = 0
338 else:
339 self.num_bad_epochs += 1
340
341 if self.num_bad_epochs > self.patience:
342 print('Tracked metric has plateaud')
343 self._reset()
344 return True
345 else:
346 return False
347
348 def is_better(self, a, best):
349
350 if self.mode == 'min' and self.threshold_mode == 'rel':
351 rel_epsilon = 1. - self.threshold
352 return a < best * rel_epsilon
353
354 elif self.mode == 'min' and self.threshold_mode == 'abs':
355 return a < best - self.threshold
356
357 elif self.mode == 'max' and self.threshold_mode == 'rel':
358 rel_epsilon = self.threshold + 1.
359 return a > best * rel_epsilon
360
361 else: # mode == 'max' and epsilon_mode == 'abs':

Callers 1

general_utils.pyFile · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected