| 4 | |
| 5 | |
| 6 | class EMA(object): |
| 7 | def __init__(self, model, decay, copy_init=False, use_double=False, inner_T=1, warmup=1): |
| 8 | self.ema_state_dict = OrderedDict() |
| 9 | self.logger = get_logger(__name__) |
| 10 | self.logger.info(f'EMA: decay={decay}, copy_init={copy_init}, \ |
| 11 | use_double={use_double}, inner_T={inner_T}, warmup={warmup}') |
| 12 | self.use_double = use_double |
| 13 | self.inner_T = inner_T |
| 14 | self.decay = decay |
| 15 | self.warmup = warmup |
| 16 | if self.inner_T > 1: |
| 17 | self.decay = self.decay ** self.inner_T |
| 18 | self.logger.info('EMA: effective decay={}'.format(self.decay)) |
| 19 | state_dict = model.state_dict() |
| 20 | if copy_init: |
| 21 | if self.use_double: |
| 22 | for k, v in state_dict.items(): |
| 23 | self.ema_state_dict[k] = v.data.clone().double() |
| 24 | else: |
| 25 | for k, v in state_dict.items(): |
| 26 | self.ema_state_dict[k] = v.data.clone().float() |
| 27 | else: |
| 28 | if self.use_double: |
| 29 | for k, v in state_dict.items(): |
| 30 | self.ema_state_dict[k] = torch.zeros_like(v).double() |
| 31 | else: |
| 32 | for k, v in state_dict.items(): |
| 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(): |
| 54 | tmp = v.data.clone() |
| 55 | v.data.copy_(self.ema_state_dict[k].data) |
| 56 | self.ema_state_dict[k].data.copy_(tmp) |
| 57 | |
| 58 | def recover(self, model): |
| 59 | state_dict = model.state_dict() |
| 60 | for k, v in self.ema_state_dict.items(): |
| 61 | tmp = v.data.clone() |
| 62 | v.data.copy_(state_dict[k].data) |
| 63 | state_dict[k].data.copy_(tmp) |
no outgoing calls
no test coverage detected