Set lr for all parameters groups.
(self, increment, token_num=None)
| 114 | return self.min_lr + coeff * delta_lr |
| 115 | |
| 116 | def step(self, increment, token_num=None): |
| 117 | """Set lr for all parameters groups.""" |
| 118 | if token_num is None: |
| 119 | args = get_args() |
| 120 | token_num = args.consumed_train_tokens |
| 121 | self.num_tokens = token_num |
| 122 | self.num_steps += increment |
| 123 | new_lr = self.get_lr() |
| 124 | for group in self.optimizer.param_groups: |
| 125 | group["lr"] = new_lr |
| 126 | |
| 127 | def state_dict(self): |
| 128 | state_dict = { |
no test coverage detected