MCPcopy Create free account
hub / github.com/MoonInTheRiver/DiffSinger / step

Method step

modules/parallel_wavegan/optimizers/radam.py:27–91  ·  view source on GitHub ↗

Run one step.

(self, closure=None)

Source from the content-addressed store, hash-verified

25 super(RAdam, self).__setstate__(state)
26
27 def step(self, closure=None):
28 """Run one step."""
29 loss = None
30 if closure is not None:
31 loss = closure()
32
33 for group in self.param_groups:
34
35 for p in group['params']:
36 if p.grad is None:
37 continue
38 grad = p.grad.data.float()
39 if grad.is_sparse:
40 raise RuntimeError('RAdam does not support sparse gradients')
41
42 p_data_fp32 = p.data.float()
43
44 state = self.state[p]
45
46 if len(state) == 0:
47 state['step'] = 0
48 state['exp_avg'] = torch.zeros_like(p_data_fp32)
49 state['exp_avg_sq'] = torch.zeros_like(p_data_fp32)
50 else:
51 state['exp_avg'] = state['exp_avg'].type_as(p_data_fp32)
52 state['exp_avg_sq'] = state['exp_avg_sq'].type_as(p_data_fp32)
53
54 exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']
55 beta1, beta2 = group['betas']
56
57 exp_avg_sq.mul_(beta2).addcmul_(1 - beta2, grad, grad)
58 exp_avg.mul_(beta1).add_(1 - beta1, grad)
59
60 state['step'] += 1
61 buffered = self.buffer[int(state['step'] % 10)]
62 if state['step'] == buffered[0]:
63 N_sma, step_size = buffered[1], buffered[2]
64 else:
65 buffered[0] = state['step']
66 beta2_t = beta2 ** state['step']
67 N_sma_max = 2 / (1 - beta2) - 1
68 N_sma = N_sma_max - 2 * state['step'] * beta2_t / (1 - beta2_t)
69 buffered[1] = N_sma
70
71 # more conservative since it's an approximated value
72 if N_sma >= 5:
73 step_size = math.sqrt(
74 (1 - beta2_t) * (N_sma - 4) / (N_sma_max - 4) * (N_sma - 2) / N_sma * N_sma_max / (N_sma_max - 2)) / (1 - beta1 ** state['step']) # NOQA
75 else:
76 step_size = 1.0 / (1 - beta1 ** state['step'])
77 buffered[2] = step_size
78
79 if group['weight_decay'] != 0:
80 p_data_fp32.add_(-group['weight_decay'] * group['lr'], p_data_fp32)
81
82 # more conservative since it's an approximated value
83 if N_sma >= 5:
84 denom = exp_avg_sq.sqrt().add_(group['eps'])

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected