(self,
params,
lr=1e-3,
bias_correction=True,
betas=(0.9, 0.999),
eps=1e-8,
adam_w_mode=True,
weight_decay=0.,
amsgrad=False,
set_grad_none=True,
ema_decay=0.9999,
use_num_upates=True
)
| 63 | """ |
| 64 | |
| 65 | def __init__(self, |
| 66 | params, |
| 67 | lr=1e-3, |
| 68 | bias_correction=True, |
| 69 | betas=(0.9, 0.999), |
| 70 | eps=1e-8, |
| 71 | adam_w_mode=True, |
| 72 | weight_decay=0., |
| 73 | amsgrad=False, |
| 74 | set_grad_none=True, |
| 75 | ema_decay=0.9999, |
| 76 | use_num_upates=True |
| 77 | ): |
| 78 | |
| 79 | if amsgrad: |
| 80 | raise RuntimeError('FusedAdam does not support the AMSGrad variant.') |
| 81 | defaults = dict(lr=lr, bias_correction=bias_correction, betas=betas, eps=eps, weight_decay=weight_decay) |
| 82 | super(FusedEmaAdam, self).__init__(params, defaults) |
| 83 | self.adam_w_mode = 1 if adam_w_mode else 0 |
| 84 | self.set_grad_none = set_grad_none |
| 85 | |
| 86 | fused_ema_adam_cuda = FusedEmaAdamBuilder().jit_load() |
| 87 | # Skip buffer |
| 88 | self._dummy_overflow_buf = get_accelerator().IntTensor([0]) |
| 89 | self.multi_tensor_ema_adam = fused_ema_adam_cuda.multi_tensor_ema_adam |
| 90 | self.ema_decay = ema_decay |
| 91 | if use_num_upates: |
| 92 | self.num_updates = 0 |
| 93 | else: |
| 94 | self.num_updates = -1 |
| 95 | self.collected_params = [] |
| 96 | |
| 97 | def zero_grad(self): |
| 98 | if self.set_grad_none: |
nothing calls this directly
no test coverage detected