(
self,
optimizer,
max_lr,
min_lr,
warmup_steps,
decay_steps,
decay_style,
use_checkpoint_lr_scheduler=True,
override_lr_scheduler=False,
)
| 24 | """Anneals the learning rate.""" |
| 25 | |
| 26 | def __init__( |
| 27 | self, |
| 28 | optimizer, |
| 29 | max_lr, |
| 30 | min_lr, |
| 31 | warmup_steps, |
| 32 | decay_steps, |
| 33 | decay_style, |
| 34 | use_checkpoint_lr_scheduler=True, |
| 35 | override_lr_scheduler=False, |
| 36 | ): |
| 37 | args = get_args() |
| 38 | # Class values. |
| 39 | self.optimizer = optimizer |
| 40 | |
| 41 | self.max_lr = float(max_lr) |
| 42 | self.min_lr = min_lr |
| 43 | assert self.min_lr >= 0.0 |
| 44 | assert self.max_lr >= self.min_lr |
| 45 | |
| 46 | self.warmup_steps = warmup_steps |
| 47 | self.num_steps = 0 |
| 48 | self.decay_steps = decay_steps |
| 49 | assert self.decay_steps > 0 |
| 50 | assert self.warmup_steps < self.decay_steps |
| 51 | |
| 52 | self.decay_tokens = args.lr_decay_tokens |
| 53 | self.num_tokens = 0 |
| 54 | self.warmup_tokens = 0 |
| 55 | |
| 56 | self.decay_style = decay_style |
| 57 | |
| 58 | self.override_lr_scheduler = override_lr_scheduler |
| 59 | self.use_checkpoint_lr_scheduler = use_checkpoint_lr_scheduler |
| 60 | if self.override_lr_scheduler: |
| 61 | assert not self.use_checkpoint_lr_scheduler, ( |
| 62 | "both override and " "use-checkpoint are set." |
| 63 | ) |
| 64 | |
| 65 | # Set the learning rate |
| 66 | self.step(0) |
| 67 | |
| 68 | print_rank_0("> learning rate decay style: {}".format(self.decay_style)) |
| 69 | |
| 70 | def get_lr(self): |
| 71 | """Learning rate decay functions from: |
nothing calls this directly
no test coverage detected