(self)
| 92 | self.lr_value = self.lr(self.step_counter) |
| 93 | |
| 94 | def get_states(self): |
| 95 | # skip DecayScheduler as it does not have persistent states |
| 96 | return {'step_counter': tensor.to_numpy(self.step_counter)[0]} |
| 97 | |
| 98 | def set_states(self, states): |
| 99 | self.step_counter = Tensor((1,)) |