| 711 | self.loss_weight_va = kwargs.get("loss_weight_va", 1.0) |
| 712 | |
| 713 | def update_ema_model(self, epoch, i_iter, len_loader): |
| 714 | if not self.flag_use_ema: |
| 715 | return |
| 716 | |
| 717 | # update teacher model with EMA |
| 718 | with torch.no_grad(): |
| 719 | if epoch > self.ema_warmup_epoch: |
| 720 | ema_decay = min( |
| 721 | 1 |
| 722 | - 1 |
| 723 | / ( |
| 724 | i_iter |
| 725 | - len_loader * self.ema_warmup_epoch |
| 726 | + 1 |
| 727 | ), |
| 728 | self.ema_decay_param, |
| 729 | ) |
| 730 | else: |
| 731 | ema_decay = 0.0 |
| 732 | |
| 733 | # update weight |
| 734 | for param_train, param_eval in zip(self.net.parameters(), self.ema_model.parameters()): |
| 735 | param_eval.data = param_eval.data * ema_decay + param_train.data * (1 - ema_decay) |
| 736 | # update bn |
| 737 | for buffer_train, buffer_eval in zip(self.net.buffers(), self.ema_model.buffers()): |
| 738 | buffer_eval.data = buffer_eval.data * ema_decay + buffer_train.data * (1 - ema_decay) |
| 739 | # buffer_eval.data = buffer_train.data |
| 740 | |
| 741 | |
| 742 | def forward(self, batch, flag_return_losses=False, flag_use_ema_infer=False, num_loop_infer=0): |