| 66 | |
| 67 | class OneCycleLR(_LRScheduler): |
| 68 | def __init__(self, optimizer, max_lr, total_steps, pct_start=0.3, anneal_strategy='cos', |
| 69 | div_factor=25.0, final_div_factor=1e4, last_epoch=-1): |
| 70 | self.max_lr = max_lr if isinstance(max_lr, list) else [max_lr] * len(optimizer.param_groups) |
| 71 | self.total_steps = total_steps |
| 72 | self.pct_start = pct_start |
| 73 | self.anneal_strategy = anneal_strategy |
| 74 | self.div_factor = div_factor |
| 75 | self.final_div_factor = final_div_factor |
| 76 | |
| 77 | self.initial_lr = [lr / self.div_factor for lr in self.max_lr] |
| 78 | self.min_lr = [lr / self.final_div_factor for lr in self.max_lr] |
| 79 | |
| 80 | super().__init__(optimizer, last_epoch) |
| 81 | |
| 82 | def get_lr(self): |
| 83 | step_num = self.last_epoch |