| 33 | self.ema_state_dict[k] = torch.zeros_like(v).float() |
| 34 | |
| 35 | def step(self, model, curr_step=None): |
| 36 | if curr_step is None: |
| 37 | decay = self.decay |
| 38 | else: |
| 39 | decay = min(self.decay, (1+curr_step)/(self.warmup+curr_step)) |
| 40 | |
| 41 | if curr_step % self.inner_T != 0: |
| 42 | return |
| 43 | |
| 44 | state_dict = model.state_dict() |
| 45 | if self.use_double: |
| 46 | for k, v in state_dict.items(): |
| 47 | self.ema_state_dict[k].mul_(decay).add_(1-decay, v.double()) |
| 48 | else: |
| 49 | for k, v in state_dict.items(): |
| 50 | self.ema_state_dict[k].mul_(decay).add_(1-decay, v.float()) |
| 51 | |
| 52 | def load_ema(self, model): |
| 53 | for k, v in model.state_dict().items(): |