| 9 | |
| 10 | class _LRScheduler(object): |
| 11 | def __init__(self, optimizer, last_iter=-1): |
| 12 | if not isinstance(optimizer, torch.optim.Optimizer) and not isinstance(optimizer, fp16.FP16_Optimizer): |
| 13 | raise TypeError('{} is not an Optimizer'.format( |
| 14 | type(optimizer).__name__)) |
| 15 | self.optimizer = optimizer |
| 16 | if last_iter == -1: |
| 17 | for group in optimizer.param_groups: |
| 18 | group.setdefault('initial_lr', group['lr']) |
| 19 | self.has_base_lrs = True |
| 20 | self._get_base_lrs_later() |
| 21 | else: |
| 22 | self.has_base_lrs = False |
| 23 | #else: |
| 24 | # for i, group in enumerate(optimizer.param_groups): |
| 25 | # if 'initial_lr' not in group: |
| 26 | # raise KeyError("param 'initial_lr' is not specified " |
| 27 | # "in param_groups[{}] when resuming an optimizer".format(i)) |
| 28 | self.last_iter = last_iter |
| 29 | |
| 30 | def _get_base_lrs_later(self): |
| 31 | self.base_lrs = list(map(lambda group: group['initial_lr'], self.optimizer.param_groups)) |