(self, zero_grad=False)
| 45 | |
| 46 | @torch.no_grad() |
| 47 | def second_step(self, zero_grad=False): |
| 48 | for group in self.param_groups: |
| 49 | for p in group["params"]: |
| 50 | if p.grad is None: continue |
| 51 | p.sub_(self.state[p]["e_w"]) # get back to "w" from "w + e(w)" |
| 52 | |
| 53 | self.base_optimizer.step() # do the actual "sharpness-aware" update |
| 54 | |
| 55 | if zero_grad: self.zero_grad() |
| 56 | |
| 57 | @torch.no_grad() |
| 58 | def step(self, closure=None): |