MCPcopy Create free account
hub / github.com/zai-org/CodeGeeX / AnnealingLR

Class AnnealingLR

codegeex/megatron/learning_rates.py:23–194  ·  view source on GitHub ↗

Anneals the learning rate.

Source from the content-addressed store, hash-verified

21
22
23class AnnealingLR(object):
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:
72 https://openreview.net/pdf?id=BJYwwY9ll pg. 4"""
73
74 # Use linear warmup for the initial part.
75 if self.warmup_steps > 0 and self.num_steps <= self.warmup_steps:
76 if self.num_steps == self.warmup_steps and self.decay_tokens is not None:
77 self.warmup_tokens = self.num_tokens
78 return self.max_lr * float(self.num_steps) / float(self.warmup_steps)
79
80 # If the learning rate is constant, just return the initial value.

Callers 1

Calls

no outgoing calls

Tested by

no test coverage detected