| 25 | DECAY_STYLES = ['linear', 'cosine', 'exponential', 'constant', 'None'] |
| 26 | |
| 27 | def __init__(self, optimizer, start_lr, warmup_iter, num_iters, decay_style=None, last_iter=-1, decay_ratio=0.5): |
| 28 | assert warmup_iter <= num_iters |
| 29 | self.optimizer = optimizer |
| 30 | self.start_lr = start_lr |
| 31 | self.warmup_iter = warmup_iter |
| 32 | self.num_iters = last_iter + 1 |
| 33 | self.end_iter = num_iters |
| 34 | self.decay_style = decay_style.lower() if isinstance(decay_style, str) else None |
| 35 | self.decay_ratio = 1 / decay_ratio |
| 36 | self.step(self.num_iters) |
| 37 | if not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0: |
| 38 | print(f'learning rate decaying style {self.decay_style}, ratio {self.decay_ratio}') |
| 39 | |
| 40 | def get_lr(self): |
| 41 | # https://openreview.net/pdf?id=BJYwwY9ll pg. 4 |