(self, params, base_optimizer, rho=0.05, **kwargs)
| 21 | |
| 22 | class SAM(torch.optim.Optimizer): |
| 23 | def __init__(self, params, base_optimizer, rho=0.05, **kwargs): |
| 24 | assert rho >= 0.0, f"Invalid rho, should be non-negative: {rho}" |
| 25 | |
| 26 | defaults = dict(rho=rho, **kwargs) |
| 27 | super(SAM, self).__init__(params, defaults) |
| 28 | |
| 29 | self.base_optimizer = base_optimizer(self.param_groups, **kwargs) |
| 30 | self.param_groups = self.base_optimizer.param_groups |
| 31 | |
| 32 | @torch.no_grad() |
| 33 | def first_step(self, zero_grad=False): |
nothing calls this directly
no outgoing calls
no test coverage detected