(self, model)
| 119 | self.registered[name] = param.clone().detach() |
| 120 | |
| 121 | def __call__(self, model): |
| 122 | self.count += 1 |
| 123 | for name, param in model.named_parameters(): |
| 124 | if param.requires_grad: |
| 125 | new_weight = param.clone().detach() if name not in self.registered else self.gamma * param + (1 - self.gamma) * self.registered[name] |
| 126 | self.registered[name] = new_weight |
| 127 | |
| 128 | if self.count % self.save_frequency == 0: |
| 129 | self.save_ema_weights() |
| 130 | |
| 131 | def copy_weights_to(self, model): |
| 132 | for name, param in model.named_parameters(): |
nothing calls this directly
no test coverage detected