| 302 | |
| 303 | |
| 304 | class 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': |