(self, param_group)
| 52 | self._add_state(param, "step", initializer=0.0) |
| 53 | |
| 54 | def _updates(self, param_group): |
| 55 | lr = param_group["lr"] |
| 56 | weight_decay = param_group["weight_decay"] |
| 57 | rho = param_group["rho"] |
| 58 | eps = param_group["eps"] |
| 59 | |
| 60 | def make_scalar(val): |
| 61 | return tensor(val, dtype="float32") |
| 62 | |
| 63 | # since `conver_inputs` is disabled for param updates, |
| 64 | # scalar should be explicitly tansforred to tensor |
| 65 | |
| 66 | _lr = make_scalar(lr) |
| 67 | _weight_decay = make_scalar(weight_decay) |
| 68 | _rho = make_scalar(rho) |
| 69 | _eps = make_scalar(eps) |
| 70 | |
| 71 | c1, c2, c05 = map(make_scalar, (1.0, 2.0, 0.5)) |
| 72 | |
| 73 | for param in param_group["params"]: |
| 74 | |
| 75 | if param.grad is None: |
| 76 | continue |
| 77 | |
| 78 | states = self._state[param] |
| 79 | step = states["step"] |
| 80 | step += c1 |
| 81 | grad = param.grad |
| 82 | if is_tracing() or weight_decay != 0.0: |
| 83 | grad = grad + param * _weight_decay |
| 84 | |
| 85 | square_avg = states["square_avg"] |
| 86 | acc_delta = states["acc_delta"] |
| 87 | square_avg = _rho * square_avg + (c1 - _rho) * grad ** c2 |
| 88 | std = (square_avg + _eps) ** c05 |
| 89 | delta = (acc_delta + _eps) ** c05 / std * grad |
| 90 | param -= _lr * delta |
| 91 | acc_delta = _rho * acc_delta + (c1 - _rho) * delta ** c2 |
| 92 | states["square_avg"]._reset(square_avg) |
| 93 | states["acc_delta"]._reset(acc_delta) |
nothing calls this directly
no test coverage detected