(self, params, lr=1e-3, betas=(0.9, 0.999), eps=1e-8,
weight_decay=1e-2, amsgrad=False, clip_norm=None, norm_type=2)
| 64 | |
| 65 | class AdamWWithClipDev(AdamW): |
| 66 | def __init__(self, params, lr=1e-3, betas=(0.9, 0.999), eps=1e-8, |
| 67 | weight_decay=1e-2, amsgrad=False, clip_norm=None, norm_type=2): |
| 68 | super(AdamWWithClipDev, self).__init__(params, lr, betas, eps, weight_decay, amsgrad) |
| 69 | self.clip_norm = clip_norm |
| 70 | self.norm_type = norm_type |
| 71 | |
| 72 | self._split_param_groups = None |
| 73 | self.reset_split_param_groups() |
| 74 | |
| 75 | def reset_split_param_groups(self): |
| 76 | if self.clip_norm is not None: |
nothing calls this directly
no test coverage detected