justyoung/DiffSinger
1
1from utils.hparams import hparams2 3 4class RSQRTSchedule(object):5 def __init__(self, optimizer):6 super().__init__()7 self.optimizer = optimizer8 self.constant_lr = hparams['lr']9 self.warmup_updates = hparams['warmup_updates']10 self.hidden_size = hparams['hidden_size']11 self.lr = hparams['lr']12 for param_group in optimizer.param_groups:13 param_group['lr'] = self.lr14 self.step(0)15 16 def step(self, num_updates):17 constant_lr = self.constant_lr18 warmup = min(num_updates / self.warmup_updates, 1.0)19 rsqrt_decay = max(self.warmup_updates, num_updates) ** -0.520 rsqrt_hidden = self.hidden_size ** -0.521 self.lr = max(constant_lr * warmup * rsqrt_decay * rsqrt_hidden, 1e-7)22 for param_group in self.optimizer.param_groups:23 param_group['lr'] = self.lr24 return self.lr25 26 def get_lr(self):27 return self.optimizer.param_groups[0]['lr']28 