(self, params, learning_rate=1e-3, beta1=0.9, beta2=0.999, eps=1e-6, weight_decay=0.0)
| 47 | """ |
| 48 | |
| 49 | def __init__(self, params, learning_rate=1e-3, beta1=0.9, beta2=0.999, eps=1e-6, weight_decay=0.0): |
| 50 | super(FP32StateAdamWeightDecay, self).__init__(params, learning_rate=learning_rate, |
| 51 | beta1=beta1, |
| 52 | beta2=beta2, |
| 53 | eps=eps, |
| 54 | weight_decay=weight_decay) |
| 55 | |
| 56 | self.moments1 = self.clone_state(self.parameters, prefix='adam_m', init='zeros') |
| 57 | self.moments2 = self.clone_state(self.parameters, prefix='adam_v', init='zeros') |
| 58 | |
| 59 | def clone_state(self, parameter_tuple, prefix, init): |
| 60 | r""" |
nothing calls this directly
no test coverage detected