AdamW that saves EMA versions of the parameters.
(self, params, lr=1.e-3, betas=(0.9, 0.999), eps=1.e-8, # TODO: check hyperparameters before using
weight_decay=1.e-2, amsgrad=False, ema_decay=0.9999, # ema decay to match previous code
ema_power=1., param_names=())
| 90 | class AdamWwithEMAandWings(optim.Optimizer): |
| 91 | # credit to https://gist.github.com/crowsonkb/65f7265353f403714fce3b2595e0b298 |
| 92 | def __init__(self, params, lr=1.e-3, betas=(0.9, 0.999), eps=1.e-8, # TODO: check hyperparameters before using |
| 93 | weight_decay=1.e-2, amsgrad=False, ema_decay=0.9999, # ema decay to match previous code |
| 94 | ema_power=1., param_names=()): |
| 95 | """AdamW that saves EMA versions of the parameters.""" |
| 96 | if not 0.0 <= lr: |
| 97 | raise ValueError("Invalid learning rate: {}".format(lr)) |
| 98 | if not 0.0 <= eps: |
| 99 | raise ValueError("Invalid epsilon value: {}".format(eps)) |
| 100 | if not 0.0 <= betas[0] < 1.0: |
| 101 | raise ValueError("Invalid beta parameter at index 0: {}".format(betas[0])) |
| 102 | if not 0.0 <= betas[1] < 1.0: |
| 103 | raise ValueError("Invalid beta parameter at index 1: {}".format(betas[1])) |
| 104 | if not 0.0 <= weight_decay: |
| 105 | raise ValueError("Invalid weight_decay value: {}".format(weight_decay)) |
| 106 | if not 0.0 <= ema_decay <= 1.0: |
| 107 | raise ValueError("Invalid ema_decay value: {}".format(ema_decay)) |
| 108 | defaults = dict(lr=lr, betas=betas, eps=eps, |
| 109 | weight_decay=weight_decay, amsgrad=amsgrad, ema_decay=ema_decay, |
| 110 | ema_power=ema_power, param_names=param_names) |
| 111 | super().__init__(params, defaults) |
| 112 | |
| 113 | def __setstate__(self, state): |
| 114 | super().__setstate__(state) |
nothing calls this directly
no outgoing calls
no test coverage detected