| 5 | |
| 6 | |
| 7 | class EMA: |
| 8 | def __init__(self, beta): |
| 9 | super().__init__() |
| 10 | self.beta = beta |
| 11 | self.step = 0 |
| 12 | |
| 13 | def update_model_average(self, ma_model, current_model): |
| 14 | for current_params, ma_params in zip(current_model.parameters(), ma_model.parameters()): |
| 15 | old_weight, up_weight = ma_params.data, current_params.data |
| 16 | ma_params.data = self.update_average(old_weight, up_weight) |
| 17 | |
| 18 | def update_average(self, old, new): |
| 19 | if old is None: |
| 20 | return new |
| 21 | return old * self.beta + (1 - self.beta) * new |
| 22 | |
| 23 | def step_ema(self, ema_model, model, step_start_ema=2000): |
| 24 | if self.step < step_start_ema: |
| 25 | self.reset_parameters(ema_model, model) |
| 26 | self.step += 1 |
| 27 | return |
| 28 | self.update_model_average(ema_model, model) |
| 29 | self.step += 1 |
| 30 | |
| 31 | def reset_parameters(self, ema_model, model): |
| 32 | ema_model.load_state_dict(model.state_dict()) |
| 33 | |
| 34 | |
| 35 | class SelfAttention(nn.Module): |